CoolFace
Modelpublic

Resilient-Coders/QnA-Safety-llama

sourceHugging Facellama3.1updated 5mo agoView on Hugging Face
0likes
handler.py84 linesDownload Raw Back to root
1import os2from typing import Any, Dict3 4import torch5from transformers import AutoModelForCausalLM, AutoTokenizer6 7BASE_MODEL = "meta-llama/Meta-Llama-3.1-8B-Instruct"8 9 10class EndpointHandler:11    def __init__(self, path: str = "") -> None:12        token = (13            os.environ.get("HF_TOKEN")14            or os.environ.get("HUGGING_FACE_HUB_TOKEN")15            or os.environ.get("HUGGINGFACE_HUB_TOKEN")16        )17        if not token:18            raise RuntimeError(19                "HF_TOKEN is not set. Add it as a secret on the Inference Endpoint "20                "so the handler can download the gated meta-llama/Meta-Llama-3.1-8B-Instruct weights."21            )22 23        tokenizer_source = path or BASE_MODEL24        self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_source)25        self.model = AutoModelForCausalLM.from_pretrained(26            BASE_MODEL,27            token=token,28            device_map="auto",29            torch_dtype=torch.bfloat16,30        )31        self.model.eval()32 33        if self.tokenizer.pad_token_id is None:34            self.tokenizer.pad_token_id = self.tokenizer.eos_token_id35 36    def __call__(self, data: Dict[str, Any]) -> Dict[str, Any]:37        inputs_payload = data.get("inputs", data)38        messages = (39            inputs_payload.get("messages")40            if isinstance(inputs_payload, dict)41            else None42        ) or data.get("messages")43 44        if not messages:45            raise ValueError(46                "Request payload must include a 'messages' list, e.g. "47                '{"inputs": {"messages": [{"role": "user", "content": "hi"}]}}.'48            )49 50        parameters: Dict[str, Any] = data.get("parameters") or {}51        max_new_tokens = int(parameters.get("max_new_tokens", 256))52        do_sample = bool(parameters.get("do_sample", False))53        temperature = float(parameters.get("temperature", 0.7))54        top_p = float(parameters.get("top_p", 0.9))55 56        inputs = self.tokenizer.apply_chat_template(57            messages,58            add_generation_prompt=True,59            tokenize=True,60            return_dict=True,61            return_tensors="pt",62        ).to(self.model.device)63 64        generate_kwargs: Dict[str, Any] = {65            "max_new_tokens": max_new_tokens,66            "do_sample": do_sample,67            "pad_token_id": self.tokenizer.pad_token_id,68            "eos_token_id": self.tokenizer.eos_token_id,69        }70        if do_sample:71            generate_kwargs["temperature"] = temperature72            generate_kwargs["top_p"] = top_p73 74        with torch.inference_mode():75            outputs = self.model.generate(**inputs, **generate_kwargs)76 77        prompt_len = inputs["input_ids"].shape[-1]78        decoded = self.tokenizer.decode(79            outputs[0][prompt_len:],80            skip_special_tokens=True,81        )82 83        return {"generated_text": decoded}84