EuroPython2022/ToxicCommentClassification
2
1import torch2import torch.nn as nn3import gradio as gr4import numpy as np5import os6import random7from transformers import AutoConfig, AutoModel, AutoTokenizer8 9 10device = torch.device('cpu')11 12 13labels = {14 0: 'toxic',15 1: 'severe_toxic',16 2: 'obscene',17 3: 'threat',18 4: 'insult',19 5: 'identity_hate',20 }21 22MODEL_NAME='roberta-base'23NUM_CLASSES=624 25MAX_LEN = 12826tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)27 28class ToxicModel(torch.nn.Module):29 def __init__(self):30 super(ToxicModel, self).__init__()31 hidden_dropout_prob: float = 0.132 layer_norm_eps: float = 1e-733 34 config = AutoConfig.from_pretrained(MODEL_NAME)35 36 config.update(37 {38 "output_hidden_states": True,39 "hidden_dropout_prob": hidden_dropout_prob,40 "layer_norm_eps": layer_norm_eps,41 "add_pooling_layer": False,42 "num_labels": NUM_CLASSES,43 }44 )45 self.transformer = AutoModel.from_pretrained(MODEL_NAME, config=config)46 self.dropout = nn.Dropout(config.hidden_dropout_prob)47 self.dropout1 = nn.Dropout(0.1)48 self.dropout2 = nn.Dropout(0.2)49 self.dropout3 = nn.Dropout(0.3)50 self.dropout4 = nn.Dropout(0.4)51 self.dropout5 = nn.Dropout(0.5)52 self.output = nn.Linear(config.hidden_size, NUM_CLASSES) 53 54 def forward(self, input_ids, attention_mask, token_type_ids):55 transformer_out = self.transformer(input_ids=input_ids, attention_mask=attention_mask, token_type_ids=token_type_ids)56 sequence_output = transformer_out[0]57 sequence_output = self.dropout(torch.mean(sequence_output, 1))58 logits1 = self.output(self.dropout1(sequence_output))59 logits2 = self.output(self.dropout2(sequence_output))60 logits3 = self.output(self.dropout3(sequence_output))61 logits4 = self.output(self.dropout4(sequence_output))62 logits5 = self.output(self.dropout5(sequence_output))63 logits = (logits1 + logits2 + logits3 + logits4 + logits5) / 564 return logits65 66 67def inference_fn(model, input_ids=None, attention_mask=None, token_type_ids=None): 68 model.eval()69 input_ids = input_ids[0].to(device) 70 attention_mask = attention_mask[0].to(device) 71 token_type_ids = token_type_ids[0].to(device) 72 73 with torch.no_grad():74 output = model(input_ids.unsqueeze(0), attention_mask.unsqueeze(0), token_type_ids.unsqueeze(0))75 out = output.sigmoid().detach().cpu().numpy().flatten()76 77 return out78 79def predict(comment=None) -> dict: 80 text = str(comment)81 text = " ".join(text.split())82 83 inputs = tokenizer.encode_plus(84 text,85 None,86 add_special_tokens=True,87 max_length=MAX_LEN,88 pad_to_max_length=True,89 return_token_type_ids=True90 )91 ids = inputs['input_ids']92 mask = inputs['attention_mask']93 token_type_ids = inputs["token_type_ids"]94 95 ids = torch.tensor(ids, dtype=torch.long),96 mask = torch.tensor(mask, dtype=torch.long),97 token_type_ids = torch.tensor(token_type_ids, dtype=torch.long),98 99 model = ToxicModel()100 101 model.load_state_dict(torch.load("toxicx_model_0.pth", map_location=torch.device(device)))102 model.to(device)103 104 predicted = inference_fn(model, ids, mask, token_type_ids)105 106 return {labels[i]: float(predicted[i]) for i in range(NUM_CLASSES)}107 108 109gr.Interface(fn=predict, 110 inputs=gr.inputs.Textbox(lines=2, placeholder="Your Comment… "),111 title="Toxic Comment Classification",112 outputs=gr.outputs.Label(num_top_classes=NUM_CLASSES)).launch()