joshcx/workers
0
1import streamlit as st2import tokenizers3from transformers import pipeline, AutoTokenizer, AutoModelForSequenceClassification4import numpy as np5import torch6import nltk7 8nltk.download("punkt")9from nltk.tokenize import sent_tokenize10 11 12class WorkerClassifier:13 def __init__(14 self, worker_model_dir, zero_shot_model_type="facebook/bart-large-mnli"15 ):16 self.zero_shot = None17 self.zero_shot_model_type = zero_shot_model_type18 self.worker_model_dir = worker_model_dir19 self.id2label = {20 0: "lauren",21 1: "betty",22 2: "doris",23 3: "hailey",24 }25 self.label2id = {v: k for k, v in self.id2label.items()}26 27 def init_models(self):28 self.ner = self.init_anonymizer()29 self.zero_shot = self.init_zero_shot()30 self.worker_model = self.init_worker_model()31 self.worker_tokenizer = self.init_worker_tokenizer()32 33 @st.cache(34 hash_funcs={35 torch.nn.parameter.Parameter: lambda _: None,36 tokenizers.Tokenizer: lambda _: None,37 tokenizers.AddedToken: lambda _: None,38 },39 allow_output_mutation=True,40 )41 def init_worker_tokenizer(self):42 return AutoTokenizer.from_pretrained(self.worker_model_dir)43 44 @st.cache(45 hash_funcs={46 torch.nn.parameter.Parameter: lambda _: None,47 tokenizers.Tokenizer: lambda _: None,48 tokenizers.AddedToken: lambda _: None,49 },50 allow_output_mutation=True,51 )52 def init_worker_model(self):53 return AutoModelForSequenceClassification.from_pretrained(54 self.worker_model_dir, problem_type="multi_label_classification"55 )56 57 def predict_worker(self, text, threshold=0.5):58 encoding = self.worker_tokenizer(text, return_tensors="pt")59 outputs = self.worker_model(**encoding)60 61 logits = outputs["logits"]62 # apply sigmoid + threshold63 sigmoid = torch.nn.Sigmoid()64 probs = sigmoid(logits.squeeze().cpu())65 predictions = np.zeros(probs.shape)66 predictions[np.where(probs >= threshold)] = 167 # turn predicted id's into actual label names68 predicted_labels = [69 [self.id2label[idx], probs[idx].detach().item()]70 for idx, label in enumerate(predictions)71 if label == 1.072 ]73 return predicted_labels74 75 @st.cache(allow_output_mutation=True)76 def init_anonymizer(self):77 return pipeline(task="ner")78 79 def anonymize(self, text: str):80 new_sentences = []81 sentences = sent_tokenize(text)82 for sent in sentences:83 result = self.ner(sent, aggregation_strategy="simple")84 for r in reversed(result):85 if r["entity_group"] == "PER":86 sent = sent[: r["start"]] + "PERSON" + sent[r["end"] :]87 new_sentences.append(sent)88 89 return " ".join(new_sentences)90 91 @st.cache(92 hash_funcs={93 tokenizers.Tokenizer: lambda _: None,94 tokenizers.AddedToken: lambda _: None,95 torch.nn.parameter.Parameter: lambda parameter: parameter.data.numpy(),96 },97 allow_output_mutation=True,98 )99 def init_zero_shot(self):100 return pipeline(101 task="zero-shot-classification", model=self.zero_shot_model_type102 )103 104 def get_personality_sentences(self, text):105 new_sentences = []106 sentences = sent_tokenize(text)107 108 for sent in sentences:109 if self.personality_sent_classifier(sent):110 new_sentences.append(sent)111 return " ".join(new_sentences)112 113 def personality_sent_classifier(self, text, threshold=0.8):114 candidate_labels = ["describing a personality trait."]115 hypothesis_template = "This example is {}"116 117 output = self.zero_shot(118 text,119 candidate_labels=candidate_labels,120 hypothesis_template=hypothesis_template,121 )122 # print(f'{text} with score {output["scores"][0]}\n')123 if output["scores"][0] > threshold:124 return True125 return False126 127 def predict(self, text):128 # first extract sentences that are relevant to personalities129 text = self.get_personality_sentences(text)130 extracted_text = text131 132 # next anonymize the sentences133 text = self.anonymize(text)134 135 # classify text136 text = self.predict_worker(text)137 return extracted_text, text138 