CoolFace
Modelpublic

samadpls/querypls-prompt2sql

sourceHugging Faceupdated 11mo agoView on Hugging Face
6likes25downloads
handler.py33 linesDownload Raw Back to root
1import torch2 3from typing import Any, Dict4from transformers import AutoModelForCausalLM, AutoTokenizer5 6 7class EndpointHandler:8    def __init__(self, path=""):9        # load model and tokenizer from path10        self.tokenizer = AutoTokenizer.from_pretrained(path)11        self.model = AutoModelForCausalLM.from_pretrained(12            path, device_map="auto", torch_dtype=torch.float16, trust_remote_code=True13        )14        self.device = "cuda" if torch.cuda.is_available() else "cpu"15 16    def __call__(self, data: Dict[str, Any]) -> Dict[str, str]:17        # process input18        inputs = data.pop("inputs", data)19        parameters = data.pop("parameters", None)20 21        # preprocess22        inputs = self.tokenizer(inputs, return_tensors="pt").to(self.device)23 24        # pass inputs with all kwargs in data25        if parameters is not None:26            outputs = self.model.generate(**inputs, **parameters)27        else:28            outputs = self.model.generate(**inputs)29 30        # postprocess the prediction31        prediction = self.tokenizer.decode(outputs[0], skip_special_tokens=True)32 33        return [{"generated_text": prediction}]