CoolFace
Modelpublic

HCKLab/BiBert-MultiTask-2

sourceHugging Facemitupdated 4y agoView on Hugging Face
1likes7downloads
handler.py85 linesDownload Raw Back to root
1from typing import Dict, List, Any2from dataclasses import dataclass3import torch4from transformers import AutoTokenizer5from transformers import pipeline6from transformers.pipelines import PIPELINE_REGISTRY7from bibert_multitask_classification import BiBert_MultiTaskPipeline8from bert_for_sequence_classification import BertForSequenceClassification9from transformers.utils import logging10from time import perf_counter11 12 13PIPELINE_REGISTRY.register_pipeline("bibert-multitask-classification", pipeline_class=BiBert_MultiTaskPipeline, pt_model=BertForSequenceClassification)14 15logging.set_verbosity_info()16logger = logging.get_logger("transformers")17 18device = torch.device("cuda" if torch.cuda.is_available() else "cpu")19 20 21@dataclass22class Task:23    id: int24    name: str25    type: str26    num_labels: int27 28tasks = [29    Task(id=0, name='label_classification', type='seq_classification', num_labels=5),30    Task(id=1, name='binary_classification', type='seq_classification', num_labels=2)31    ]32 33 34idtolabel = {"0":"Negative", "1":"Negative", "2": "Neutral",  "3":"Positive", "4": "Positive" }35idtoscore =  {"0": -1, "1": -1, "2": 0,  "3": 1, "4": 1 }36 37class EndpointHandler():38    def __init__(self, path=""):39        # Preload all the elements you are going to need at inference.40        logger.info("The device is %s.", device)41 42        t0 = perf_counter()43 44        tokenizer = AutoTokenizer.from_pretrained(path)45        model = BertForSequenceClassification.from_pretrained(path, tasks_map=tasks).to(device)46        self.classifier_s = pipeline("bibert-multitask-classification", model = model, task_id="0", tokenizer=tokenizer, device = device)47        self.classifier_p = pipeline("bibert-multitask-classification", model = model, task_id="1", tokenizer=tokenizer, device = device)48        elapsed = 1000 * (perf_counter() - t0)49        logger.info("Models and tokenizer Polarity loaded in %d ms.", elapsed)50 51 52    def __call__(self, data: Dict[str, Any]) -> List[Dict[str, Any]]:53        """54        data args:55            inputs (:obj: `str` | `PIL.Image` | `np.array`)56            kwargs57        Return:58            A :obj:`list` | `dict`: will be serialized and returned59        """60    61        inputs = data.pop("inputs", data)62        #lang = data.pop("lang", None)63        #logger.info("The language of Verbatim is %s.", lang)64        if isinstance(inputs, str):65            inputs = [inputs]66    67        t0 = perf_counter()68        prediction_res = []69        classifier_pol = self.classifier_p(inputs)70        classifier_subj = self.classifier_s(inputs)71        logger.info("Prediction polarity %s", classifier_pol)72        logger.info("Prediction subjective %s", classifier_subj)73 74        for idx, x in enumerate(inputs):75            label = classifier_pol[idx]['label']76            prob = classifier_pol[idx]['probability']77 78            if label == '0' and prob >= 0.75:79                prediction_res.append({"label":"Neutral", "score":0}) 80            else: 81                prediction_res.append({"label":idtolabel.get(classifier_subj[idx]['label']), "score": idtoscore.get(classifier_subj[idx]['label'])})82        elapsed = 1000 * (perf_counter() - t0)83        logger.info("Model prediction time: %d ms.", elapsed)  84        return prediction_res85