INNER CODE UNIT · Python
mask
JackHCC/Chinese-Text-Classification-PyTorch · pretrain_predict.py:37
mask = [1] * len(token_ids) + ([0] * (pad_size - len(token)))
token_ids += ([0] * (pad_size - len(token)))
else:
mask = [1] * pad_size
token_ids = token_ids[:pad_size]
seq_len = pad_size
ids = torch.LongTensor([token_ids])
seq_len = torch.LongTensor([seq_len])
mask = torch.LongTensor([mask])
return ids, seq_len, mask
def predict(self, query):
# 返回预测的索引
data = self.build_predict_text(query)
with torch.no_grad():
outputs = self.model(data)
num = torch.argmax(outputs)
return key[int(num)]