CoolFace
Apppublic

joshcx/workers

sourceHugging Faceupdated 4y agoView on Hugging Face
0likes
worker.py138 linesDownload Raw Back to root
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