Sp1rit/transformer_devops
0
1import torch2import torch.nn as nn3import numpy as np4import streamlit as st5from transformers import DistilBertModel, DistilBertTokenizerFast6 7 8TARGET_IND2LABEL = {9 0: 'Computer Science',10 1: 'Economics',11 2: 'Electrical Engineering and Systems Science',12 3: 'Mathematics',13 4: 'Physics',14 5: 'Quantitative Biology',15 6: 'Quantitative Finance',16 7: 'Statistics',17}18 19class DistilBERTClassifier(nn.Module):20 def __init__(self, num_classes=8):21 super().__init__()22 self.encoder = DistilBertModel.from_pretrained("distilbert-base-cased")23 self.pre_classifier = nn.Linear(768, 768)24 self.gelu = nn.GELU()25 self.dropout = nn.Dropout(0.1)26 self.classifier = nn.Linear(768, num_classes)27 28 def forward(self, input_ids, attention_mask, labels):29 output = self.encoder(input_ids=input_ids, attention_mask=attention_mask)30 hidden_state = output[0]31 pooler = hidden_state[:, 0]32 pooler = self.dropout(self.gelu(self.pre_classifier(pooler)))33 preds = self.classifier(pooler)34 return preds35 36@st.cache_resource37def load_tokenizer():38 return DistilBertTokenizerFast.from_pretrained('distilbert-base-cased')39 40@st.cache_resource41def load_model(device):42 model = torch.load('model.pt', map_location=torch.device('cpu')).to(device)43 model.eval()44 return model45 46def get_verdict(preds):47 inds = np.argsort(preds)[::-1]48 sum_prob = 0.049 verdict = []50 for ind in inds:51 prob = preds[ind]52 sum_prob += prob53 verdict.append(f"{TARGET_IND2LABEL[ind]}: {prob}")54 if (sum_prob >= 0.95):55 break56 return "\n\n".join(verdict)57 58def get_preds(text, model, tokenizer, device):59 tokens = tokenizer(text, padding=True, truncation=True, return_tensors='pt')60 tokens['input_ids'] = tokens['input_ids'].to(device)61 tokens['attention_mask'] = tokens['attention_mask'].to(device)62 tokens['labels'] = None # made for training convinience63 with torch.no_grad():64 preds = torch.softmax(model(**tokens)[0], 0).cpu().numpy()65 return preds66 