CoolFace
Apppublic

kol/Text_Classification

sourceHugging Faceafl-3.0updated 4y agoView on Hugging Face
0likes
processing.py38 linesDownload Raw Back to root
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