Team Ai
Modelpublic

ericaRC/example

sourceHugging Facecc-by-nc-4.0updated 6mo agoView on Hugging Face
0likes5downloads
handler.py97 linesDownload Raw Back to root
1"""Custom inference handler for Hugging Face Inference Endpoints.2 3NLLB needs a source-language code on the tokenizer and a forced BOS token4id for the target language at generation time, so the default translation5pipeline is not flexible enough. This handler accepts `src_lang` and6`tgt_lang` (NLLB Flores-200 codes, e.g. "eng_Latn", "spa_Latn") per7request.8 9Request format:10    {11      "inputs": "Hello, world!",            # str or List[str]12      "parameters": {13        "src_lang": "eng_Latn",             # optional, default eng_Latn14        "tgt_lang": "spa_Latn",             # optional, default spa_Latn15        "max_length": 256,                  # optional16        "num_beams": 4,                     # optional17        "temperature": 1.0,                 # optional18        "do_sample": false                  # optional19      }20    }21 22Response: List[{"translation_text": str}]23"""24 25from __future__ import annotations26 27from typing import Any, Dict, List, Union28 29import torch30from transformers import AutoModelForSeq2SeqLM, AutoTokenizer31 32DEFAULT_SRC_LANG = "eng_Latn"33DEFAULT_TGT_LANG = "spa_Latn"34DEFAULT_MAX_LENGTH = 25635DEFAULT_NUM_BEAMS = 436 37 38class EndpointHandler:39    def __init__(self, path: str = "") -> None:40        self.device = "cuda" if torch.cuda.is_available() else "cpu"41        # fp16 on GPU keeps latency and memory down; stay in fp32 on CPU for stability.42        dtype = torch.float16 if self.device == "cuda" else torch.float3243 44        self.tokenizer = AutoTokenizer.from_pretrained(path)45        self.model = AutoModelForSeq2SeqLM.from_pretrained(46            path, torch_dtype=dtype47        ).to(self.device)48        self.model.eval()49 50    def __call__(51        self, data: Dict[str, Any]52    ) -> List[Dict[str, str]]:53        inputs: Union[str, List[str], None] = data.get("inputs")54        if inputs is None:55            return [{"error": "Missing 'inputs' field."}]56        if isinstance(inputs, str):57            inputs = [inputs]58        if not all(isinstance(x, str) for x in inputs):59            return [{"error": "'inputs' must be a string or a list of strings."}]60 61        params: Dict[str, Any] = data.get("parameters") or {}62        src_lang = params.get("src_lang", DEFAULT_SRC_LANG)63        tgt_lang = params.get("tgt_lang", DEFAULT_TGT_LANG)64        max_length = int(params.get("max_length", DEFAULT_MAX_LENGTH))65        num_beams = int(params.get("num_beams", DEFAULT_NUM_BEAMS))66        do_sample = bool(params.get("do_sample", False))67        temperature = float(params.get("temperature", 1.0))68 69        try:70            forced_bos_token_id = self.tokenizer.convert_tokens_to_ids(tgt_lang)71        except Exception:72            return [{"error": f"Unknown target language code: {tgt_lang!r}"}]73        if forced_bos_token_id == self.tokenizer.unk_token_id:74            return [{"error": f"Unknown target language code: {tgt_lang!r}"}]75 76        self.tokenizer.src_lang = src_lang77        encoded = self.tokenizer(78            inputs,79            return_tensors="pt",80            padding=True,81            truncation=True,82            max_length=max_length,83        ).to(self.device)84 85        with torch.inference_mode():86            generated = self.model.generate(87                **encoded,88                forced_bos_token_id=forced_bos_token_id,89                max_length=max_length,90                num_beams=num_beams,91                do_sample=do_sample,92                temperature=temperature,93            )94 95        decoded = self.tokenizer.batch_decode(generated, skip_special_tokens=True)96        return [{"translation_text": t} for t in decoded]97