ericaRC/example
05
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 