# handler.py import os import torch from typing import Dict, Any, Optional from transformers import ( AutoTokenizer, AutoModelForCausalLM, GenerationConfig, ) class EndpointHandler: def __init__(self, model_path: str = None): # Load tokenizer & model self.tokenizer = AutoTokenizer.from_pretrained(model_path, use_fast=True) self.model = AutoModelForCausalLM.from_pretrained(model_path) # Ensure pad token if self.tokenizer.pad_token_id is None: self.tokenizer.pad_token = self.tokenizer.eos_token # Put model on GPU if available self.device = "cuda" if torch.cuda.is_available() else "cpu" self.model.to(self.device) def __call__(self, data: Dict[str, Any]) -> Dict[str, Any]: # 1) Extract user input & parameters user_input = data.get("inputs", "") params = data.get("parameters", {}) # 2) Build prompt prompt = f"User: {user_input}\nAssistant:" # 3) Tokenize & prepare tensors encoded = self.tokenizer( prompt, return_tensors="pt", padding=False, ).to(self.device) input_ids = encoded.input_ids attention_mask = encoded.attention_mask prompt_len = input_ids.shape[1] # 4) Merge default gen args with overrides gen_args = { "max_new_tokens": params.get("max_new_tokens", 128), "temperature": params.get("temperature", 1.0), "top_p": params.get("top_p", 1.0), "top_k": params.get("top_k", 50), "do_sample": params.get("do_sample", True), "repetition_penalty": params.get("repetition_penalty", 1.0), "pad_token_id": self.tokenizer.pad_token_id, "eos_token_id": self.tokenizer.eos_token_id, } # Build a GenerationConfig for cleaner API gen_config = GenerationConfig(**gen_args) # 5) Call generate output = self.model.generate( inputs=input_ids, attention_mask=attention_mask, generation_config=gen_config, ) # 6) Decode only the new tokens # output is shape [1, prompt_len + gen_len] full = self.tokenizer.decode(output[0, prompt_len:], skip_special_tokens=True) first_reply = full.split("\nUser:")[0].strip() return {"generated_text": first_reply}