Views
No views yet
1 def get_prediction(text):
2 encoding = tokenizer(text, return_tensors="pt", padding="max_length", truncation=True, max_length=128)
3 encoding = {k: v.to(trainer.model.device) for k,v in encoding.items()}
4
5 outputs = model(**encoding)
6
7 logits = outputs.logits
8
9 sigmoid = torch.nn.Sigmoid()
10 probs = sigmoid(logits.squeeze().cpu())
11 probs = probs.detach().numpy()
12 label = np.argmax(probs, axis=-1)
13 if label == 1:
14 if probs[1] > 0.7:
15 return 1
16 else:
17 return 0
18 else:
19 return 0