import torch import torch.nn as nn def MakePrediction(model, model_head, tokenizer, title, summary=None): classes_list = ["computer science", "math", "biology", "economy", "statistics", "physics"] text = title if summary: text += summary text_info = tokenizer(text, truncation=True, return_tensors="pt", padding=True) # text_info = {k: v.to(device) for k, v in text_info.items()} with torch.no_grad(): ans = model(**text_info) ans = ans.last_hidden_state[:, 0] ans = model_head(ans) # sigm = nn.Sigmoid() probs = nn.Softmax() # ans = sigm(ans) ans = probs(ans) answers_idx = torch.cat((ans.view(6,1), torch.arange(6)[:,None]), 1).tolist() answers_idx.sort(reverse=True, key=lambda x: x[0]) classes_idx = [int(answers_idx[0][1])] probs = [answers_idx[0][0]] summ_prob = probs[0] for i in range(1, 6): if summ_prob > 0.95: break summ_prob += answers_idx[i][0] probs.append(answers_idx[i][0]) classes_idx.append(answers_idx[i][1]) classes = [] for i in classes_idx: classes.append(classes_list[int(i)]) return classes, probs