Ephraimmm's picture
Update handler.py
5d7d3c3 verified
Raw
History Blame Contribute Delete
5.74 kB
from transformers import AutoModelForCausalLM, AutoTokenizer
import torch
class EndpointHandler:
def __init__(self, path: str = ""):
# Load tokenizer
self.tokenizer = AutoTokenizer.from_pretrained(
path,
trust_remote_code=True,
use_auth_token=True,
)
if self.tokenizer.pad_token_id is None:
self.tokenizer.pad_token = self.tokenizer.eos_token
# Load model
self.model = AutoModelForCausalLM.from_pretrained(
path,
torch_dtype="auto",
device_map="auto",
trust_remote_code=True,
use_auth_token=True,
)
self.model.eval()
print("✓ Model loaded successfully")
def __call__(self, data):
prompt = data["inputs"]
inputs = self.tokenizer(prompt, return_tensors="pt")
inputs = {k: v.to(self.model.device) for k, v in inputs.items()}
with torch.inference_mode():
outputs = self.model.generate(
**inputs,
max_new_tokens=128,
do_sample=True,
temperature=0.7,
top_p=0.9,
pad_token_id=self.tokenizer.pad_token_id,
)
text = self.tokenizer.decode(outputs[0], skip_special_tokens=True)
return [{"generated_text": text}]
# from typing import Dict, List, Any
# import torch
# from transformers import AutoModelForCausalLM, AutoTokenizer
# class EndpointHandler:
# """
# Custom handler for HuggingFace Inference Endpoints
# Handles Nigerian Pidgin English text generation
# """
# def __init__(self, path: str = ""):
# # Load tokenizer first (safer for remote-code models)
# self.tokenizer = AutoTokenizer.from_pretrained(
# path,
# trust_remote_code=True,
# use_fast=True,
# )
# # Some tokenizers have no pad token; align to eos to avoid generate() errors
# if self.tokenizer.pad_token_id is None:
# self.tokenizer.pad_token = self.tokenizer.eos_token
# # Load model
# self.model = AutoModelForCausalLM.from_pretrained(
# path,
# torch_dtype="auto",
# device_map="auto",
# trust_remote_code=True,
# )
# self.model.eval()
# self.default_system_prompt = (
# "You are a helpful assistant that speaks Nigerian Pidgin English. "
# "Respond naturally in Pidgin."
# )
# # Pick a stable device for inputs (first shard device if sharded)
# self._device = next(iter(self.model.hf_device_map.values()))
# if isinstance(self._device, str) and self._device.startswith("cuda"):
# self._device = torch.device(self._device)
# elif self._device == "cpu":
# self._device = torch.device("cpu")
# print("✓ Model and tokenizer loaded successfully")
# def __call__(self, data: Dict[str, Any]) -> List[Dict[str, Any]]:
# inputs_text = data.get("inputs", data)
# parameters = data.get("parameters", {}) or {}
# system_prompt = parameters.get("system_prompt", self.default_system_prompt)
# max_new_tokens = int(parameters.get("max_new_tokens", 100))
# temperature = float(parameters.get("temperature", 0.7))
# top_p = float(parameters.get("top_p", 0.9))
# top_k = int(parameters.get("top_k", 50))
# repetition_penalty = float(parameters.get("repetition_penalty", 1.1))
# do_sample = bool(parameters.get("do_sample", True))
# return_full_text = bool(parameters.get("return_full_text", False))
# # Prefer chat template if tokenizer supports it
# if hasattr(self.tokenizer, "apply_chat_template"):
# messages = []
# if system_prompt:
# messages.append({"role": "system", "content": system_prompt})
# messages.append({"role": "user", "content": str(inputs_text)})
# prompt = self.tokenizer.apply_chat_template(
# messages,
# tokenize=False,
# add_generation_prompt=True,
# )
# else:
# # Fallback
# if system_prompt:
# prompt = f"{system_prompt}\n\nUser: {inputs_text}\nAssistant:"
# else:
# prompt = str(inputs_text)
# enc = self.tokenizer(
# prompt,
# return_tensors="pt",
# truncation=True,
# max_length=2048,
# )
# # Move only input tensors to the chosen device
# enc = {k: v.to(self._device) for k, v in enc.items()}
# with torch.inference_mode():
# out = self.model.generate(
# **enc,
# max_new_tokens=max_new_tokens,
# do_sample=do_sample,
# temperature=temperature if do_sample else None,
# top_p=top_p if do_sample else None,
# top_k=top_k if do_sample else None,
# repetition_penalty=repetition_penalty,
# pad_token_id=self.tokenizer.pad_token_id,
# eos_token_id=self.tokenizer.eos_token_id,
# )
# decoded = self.tokenizer.decode(out[0], skip_special_tokens=True)
# if not return_full_text:
# # If we used chat template, easiest is to strip the prompt prefix
# if decoded.startswith(prompt):
# decoded = decoded[len(prompt):].strip()
# elif "Assistant:" in decoded:
# decoded = decoded.split("Assistant:")[-1].strip()
# return [{"generated_text": decoded}]