joelb/custom-handler-tutorial
116
1from typing import Dict, List, Any2from transformers import pipeline3import holidays4 5 6class EndpointHandler:7 def __init__(self, path=""):8 self.pipeline = pipeline("text-classification", model=path)9 self.holidays = holidays.US()10 11 def __call__(self, data: Dict[str, Any]) -> List[Dict[str, Any]]:12 """13 data args:14 inputs (:obj: `str`)15 date (:obj: `str`)16 Return:17 A :obj:`list` | `dict`: will be serialized and returned18 """19 # get inputs20 inputs = data.pop("inputs", data)21 # get additional date field22 date = data.pop("date", None)23 24 # check if date exists and if it is a holiday25 if date is not None and date in self.holidays:26 return [{"label": "happy", "score": 1}]27 28 # run normal prediction29 prediction = self.pipeline(inputs)30 return prediction31 