kol/Text_Classification
0
1import torch2import torch.nn as nn3 4def MakePrediction(model, model_head, tokenizer, title, summary=None):5 classes_list = ["computer science", "math", "biology", "economy", "statistics", "physics"]6 7 text = title8 if summary:9 text += summary10 text_info = tokenizer(text, truncation=True, return_tensors="pt", padding=True)11# text_info = {k: v.to(device) for k, v in text_info.items()}12 with torch.no_grad():13 ans = model(**text_info)14 ans = ans.last_hidden_state[:, 0]15 ans = model_head(ans)16# sigm = nn.Sigmoid()17 probs = nn.Softmax()18 # ans = sigm(ans)19 ans = probs(ans)20 21 answers_idx = torch.cat((ans.view(6,1), torch.arange(6)[:,None]), 1).tolist()22 answers_idx.sort(reverse=True, key=lambda x: x[0])23 24 classes_idx = [int(answers_idx[0][1])]25 probs = [answers_idx[0][0]]26 summ_prob = probs[0]27 28 for i in range(1, 6):29 if summ_prob > 0.95:30 break31 summ_prob += answers_idx[i][0]32 probs.append(answers_idx[i][0])33 classes_idx.append(answers_idx[i][1])34 35 classes = []36 for i in classes_idx:37 classes.append(classes_list[int(i)])38 return classes, probs