PIDGIN_gemma3 / handler.py
Ephraimmm's picture
Update handler.py
52ad975 verified
Raw
History Blame Contribute Delete
17.8 kB
#tope_version
# handler.py
# handler.py
# from typing import Any, Dict, List, Union, Optional
# import torch
# from huggingface_inference_toolkit.logging import logger
# from unsloth import FastLanguageModel
# from peft import PeftModel
# # OpenAI-style prompt:
# # - str
# # - [{"role": "...", "content": "..." | [{"type":"text","text":"..."}]}]
# Prompt = Union[str, List[Dict[str, Any]], Dict[str, Any]]
# def _to_block_list(content: Any) -> List[Dict[str, str]]:
# """Ensure content is a list[{'type':'text','text':...}]."""
# if content is None:
# return [{"type": "text", "text": ""}]
# if isinstance(content, list):
# # Assume already blocks
# return content
# # Fallback: stringify
# return [{"type": "text", "text": str(content)}]
# class EndpointHandler:
# """
# HF Inference Endpoint handler that mirrors your Gradio setup:
# - Loads Unsloth 4-bit Gemma-3-4B IT
# - Attaches LoRA adapter
# - Uses tokenizer.apply_chat_template
# - Always replies in Nigerian Pidgin
# """
# def __init__(self, model_dir: str, **_: Any):
# logger.info(f"Initializing with model_dir={model_dir}")
# base_model_name = "unsloth/gemma-3-4b-it-unsloth-bnb-4bit"
# model, tokenizer = FastLanguageModel.from_pretrained(
# model_name=base_model_name,
# max_seq_length=2048,
# dtype=torch.float16,
# load_in_4bit=True,
# )
# lora_repo = "Ephraimmm/PIDGIN_gemma-3"
# model = PeftModel.from_pretrained(model, lora_repo)
# FastLanguageModel.for_inference(model)
# self.model = model.eval()
# self.tokenizer = tokenizer
# self.device = getattr(
# self.model,
# "device",
# torch.device("cuda" if torch.cuda.is_available() else "cpu"),
# )
# # Safety: some tokenizers lack pad_token_id
# if getattr(self.tokenizer, "pad_token_id", None) is None:
# self.tokenizer.pad_token_id = self.tokenizer.eos_token_id
# logger.info(f"Device: {self.device} | eos_id: {self.tokenizer.eos_token_id}")
# self._system_text = (
# "You are a Nigerian assistant that speaks PIDGIN ENGLISH. "
# "When asked 'how far', reply 'I dey o, how you dey?'. "
# "Always answer in Pidgin English."
# )
# def _normalize_messages(self, prompt: Prompt) -> List[Dict[str, Any]]:
# """Return a list of messages with block-style contents only."""
# # Accept single dict
# if isinstance(prompt, dict):
# prompt = [prompt]
# # Accept raw string
# if isinstance(prompt, str):
# msgs = [
# {"role": "system", "content": _to_block_list(self._system_text)},
# {"role": "user", "content": _to_block_list(prompt)},
# ]
# return msgs
# if not isinstance(prompt, list):
# raise ValueError("`inputs` must be a string, a message dict, or a list of messages.")
# # Normalize all contents to blocks
# norm = []
# for m in prompt:
# role = m.get("role", "user")
# content = _to_block_list(m.get("content", ""))
# norm.append({"role": role, "content": content})
# # Ensure a system message exists (prepend if missing)
# if not any(m.get("role") == "system" for m in norm):
# norm = [{"role": "system", "content": _to_block_list(self._system_text)}] + norm
# return norm
# def _build_model_inputs(self, prompt: Prompt):
# messages = self._normalize_messages(prompt)
# # Use the same path as your Gradio script
# model_inputs = self.tokenizer.apply_chat_template(
# messages,
# add_generation_prompt=True,
# return_tensors="pt",
# tokenize=True,
# return_dict=True,
# ).to(self.device)
# return model_inputs
# def __call__(self, data: Dict[str, Any]) -> Dict[str, Any]:
# logger.info(f"Incoming keys: {list(data.keys())}")
# if "inputs" not in data:
# raise ValueError("Missing `inputs` in request body.")
# prompt: Prompt = data["inputs"]
# params: Dict[str, Any] = (data.get("parameters") or {})
# # Defaults aligned with your Gradio code
# max_new_tokens = int(params.get("max_new_tokens", 256)) # raise per request if needed
# temperature = float(params.get("temperature", 0.1))
# top_p = float(params.get("top_p", 1.0))
# top_k: Optional[int] = params.get("top_k", None)
# use_cache = bool(params.get("use_cache", False))
# repetition_penalty = float(params.get("repetition_penalty", 1.0))
# eos_token_id = params.get("eos_token_id", self.tokenizer.eos_token_id)
# pad_token_id = params.get("pad_token_id", self.tokenizer.pad_token_id)
# inputs = self._build_model_inputs(prompt)
# gen_kwargs = dict(
# **inputs,
# max_new_tokens=max_new_tokens,
# temperature=temperature,
# top_p=top_p,
# do_sample=True, # streaming-like sampling on
# use_cache=use_cache,
# repetition_penalty=repetition_penalty,
# eos_token_id=eos_token_id,
# pad_token_id=pad_token_id,
# )
# if top_k is not None:
# gen_kwargs["top_k"] = int(top_k)
# with torch.no_grad():
# output_ids = self.model.generate(**gen_kwargs)
# # Return ONLY the continuation (like your streamer)
# input_len = inputs["input_ids"].shape[-1]
# generated_ids = output_ids[0, input_len:]
# reply_text = self.tokenizer.decode(generated_ids, skip_special_tokens=True)
# return {
# "reply": [reply_text],
# "usage": {
# "prompt_tokens": int(input_len),
# "generated_tokens": int(generated_ids.shape[-1]),
# },
# }
# if __name__ == "__main__":
# h = EndpointHandler(model_dir=".")
# # Raw string
# print(h({"inputs": "How far?", "parameters": {"max_new_tokens": 64}}))
# # Block-style content (your Gradio shape)
# print(h({
# "inputs": [
# {
# "role": "user",
# "content": [{"type": "text", "text": "Explain why sky dey blue."}]
# }
# ],
# "parameters": {"max_new_tokens": 64}
# }))
from typing import Any, Dict, List, Union, Optional
import torch
import time
from huggingface_inference_toolkit.logging import logger
from unsloth import FastLanguageModel
from peft import PeftModel
# OpenAI-style prompt:
# - str
# - [{"role": "...", "content": "..." | [{"type":"text","text":"..."}]}]
Prompt = Union[str, List[Dict[str, Any]], Dict[str, Any]]
def _to_block_list(content: Any) -> List[Dict[str, str]]:
"""Ensure content is a list[{'type':'text','text':...}]."""
if content is None:
return [{"type": "text", "text": ""}]
if isinstance(content, list):
# Check if it's already in block format
if all(isinstance(item, dict) and "type" in item for item in content):
return content
# Otherwise treat as raw list, stringify
return [{"type": "text", "text": str(content)}]
# Fallback: stringify
return [{"type": "text", "text": str(content)}]
class EndpointHandler:
"""
HF Inference Endpoint handler that accepts pure OpenAI API format:
- Accepts OpenAI-style requests with 'messages', 'model', etc.
- Returns OpenAI-style responses with 'choices', 'usage', etc.
- Loads Unsloth 4-bit Gemma-3-4B IT
- Attaches LoRA adapter
- Uses tokenizer.apply_chat_template
- Always replies in Nigerian Pidgin
"""
def __init__(self, model_dir: str, **_: Any):
logger.info(f"Initializing with model_dir={model_dir}")
base_model_name = "unsloth/gemma-3-4b-it-unsloth-bnb-4bit"
model, tokenizer = FastLanguageModel.from_pretrained(
model_name=base_model_name,
max_seq_length=2048,
dtype=torch.float16,
load_in_4bit=True,
)
lora_repo = "Ephraimmm/PIDGIN_gemma-3"
model = PeftModel.from_pretrained(model, lora_repo)
FastLanguageModel.for_inference(model)
self.model = model.eval()
self.tokenizer = tokenizer
self.device = getattr(
self.model,
"device",
torch.device("cuda" if torch.cuda.is_available() else "cpu"),
)
# Safety: some tokenizers lack pad_token_id
if getattr(self.tokenizer, "pad_token_id", None) is None:
self.tokenizer.pad_token_id = self.tokenizer.eos_token_id
logger.info(f"Device: {self.device} | eos_id: {self.tokenizer.eos_token_id}")
self._system_text = (
"You are a Nigerian assistant that speaks PIDGIN ENGLISH. "
"When asked 'how far', reply 'I dey o, how you dey?'. "
"Always answer in Pidgin English."
)
def _normalize_messages(self, messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
"""Normalize OpenAI-style messages to block-style contents."""
norm = []
for m in messages:
role = m.get("role", "user")
content = m.get("content", "")
# Convert content to block list format for internal processing
if isinstance(content, str):
content = _to_block_list(content)
elif isinstance(content, list):
content = _to_block_list(content)
else:
content = _to_block_list(str(content))
norm.append({"role": role, "content": content})
# Ensure a system message exists (prepend if missing)
if not any(m.get("role") == "system" for m in norm):
norm = [{"role": "system", "content": _to_block_list(self._system_text)}] + norm
return norm
def _build_model_inputs(self, messages: List[Dict[str, Any]]):
normalized_messages = self._normalize_messages(messages)
model_inputs = self.tokenizer.apply_chat_template(
normalized_messages,
add_generation_prompt=True,
return_tensors="pt",
tokenize=True,
return_dict=True,
).to(self.device)
return model_inputs
def __call__(self, data: Dict[str, Any]) -> Dict[str, Any]:
"""
Process requests in either OpenAI or HF format and return OpenAI-style responses.
Input format (OpenAI API - directly from client):
{
"model": "Ephraimmm/PIDGIN_gemma-3",
"messages": [
{"role": "user", "content": "What is deep learning?"}
],
"temperature": 0.7,
"max_tokens": 256
}
OR HF Inference Endpoint format (with inputs wrapper):
{
"inputs": {
"model": "Ephraimmm/PIDGIN_gemma-3",
"messages": [{"role": "user", "content": "What is deep learning?"}]
}
}
Output format (OpenAI API):
{
"id": "chatcmpl-...",
"object": "chat.completion",
"created": 1234567890,
"model": "Ephraimmm/PIDGIN_gemma-3",
"choices": [{
"index": 0,
"message": {
"role": "assistant",
"content": "Response text"
},
"finish_reason": "stop"
}],
"usage": {
"prompt_tokens": 10,
"completion_tokens": 20,
"total_tokens": 30
}
}
"""
logger.info(f"Incoming request keys: {list(data.keys())}")
logger.info(f"Full incoming data: {data}")
# Handle different input formats
# Case 1: HF sends {"inputs": "string"} - treat as user message
if "inputs" in data and isinstance(data["inputs"], str):
messages = [{"role": "user", "content": data["inputs"]}]
model_name = "Ephraimmm/PIDGIN_gemma-3"
params = data.get("parameters", {})
# Case 2: HF sends {"inputs": {...OpenAI format...}}
elif "inputs" in data and isinstance(data["inputs"], dict):
openai_data = data["inputs"]
# Preserve any parameters at root level
if "parameters" in data:
for key, value in data["parameters"].items():
if key not in openai_data:
openai_data[key] = value
messages = openai_data.get("messages")
model_name = openai_data.get("model", "Ephraimmm/PIDGIN_gemma-3")
params = openai_data
# Case 3: HF sends {"inputs": [messages array]}
elif "inputs" in data and isinstance(data["inputs"], list):
messages = data["inputs"]
model_name = "Ephraimmm/PIDGIN_gemma-3"
params = data.get("parameters", {})
# Case 4: Direct OpenAI format {"messages": [...]}
elif "messages" in data:
messages = data["messages"]
model_name = data.get("model", "Ephraimmm/PIDGIN_gemma-3")
params = data
else:
raise ValueError("Missing required field: 'messages' or 'inputs'")
# Validate messages
if not messages:
raise ValueError("'messages' cannot be empty")
messages = data["messages"]
if not isinstance(messages, list):
raise ValueError("'messages' must be a list of message objects")
messages = data["messages"]
if not isinstance(messages, list):
raise ValueError("'messages' must be a list of message objects")
# Extract OpenAI-style parameters with defaults
max_tokens = int(params.get("max_tokens", 256))
temperature = float(params.get("temperature", 0.1))
top_p = float(params.get("top_p", 1.0))
top_k = params.get("top_k", None)
stream = params.get("stream", False) # Not implemented yet
presence_penalty = params.get("presence_penalty", 0.0)
frequency_penalty = params.get("frequency_penalty", 0.0)
# Build model inputs
inputs = self._build_model_inputs(messages)
# Prepare generation parameters
gen_kwargs = dict(
**inputs,
max_new_tokens=max_tokens,
temperature=temperature,
top_p=top_p,
do_sample=temperature > 0, # Use sampling if temperature > 0
use_cache=True,
repetition_penalty=1.0 + frequency_penalty, # Map frequency_penalty
eos_token_id=self.tokenizer.eos_token_id,
pad_token_id=self.tokenizer.pad_token_id,
)
if top_k is not None:
gen_kwargs["top_k"] = int(top_k)
# Generate response
with torch.no_grad():
output_ids = self.model.generate(**gen_kwargs)
# Extract only the generated tokens (exclude prompt)
input_len = inputs["input_ids"].shape[-1]
generated_ids = output_ids[0, input_len:]
reply_text = self.tokenizer.decode(generated_ids, skip_special_tokens=True)
# Calculate token usage
prompt_tokens = int(input_len)
completion_tokens = int(generated_ids.shape[-1])
total_tokens = prompt_tokens + completion_tokens
# Return OpenAI-compatible response
response = {
"id": f"chatcmpl-{int(time.time() * 1000)}",
"object": "chat.completion",
"created": int(time.time()),
"model": model_name,
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": reply_text
},
"finish_reason": "stop"
}
],
"usage": {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": total_tokens
}
}
return response
if __name__ == "__main__":
h = EndpointHandler(model_dir=".")
# Test 1: OpenAI format (direct)
print("=== Test 1: OpenAI Format (Direct) ===")
response = h({
"model": "Ephraimmm/PIDGIN_gemma-3",
"messages": [
{"role": "user", "content": "What is deep learning?"}
]
})
print(response)
print()
# Test 2: HF Inference Endpoint format (with inputs wrapper)
print("=== Test 2: HF Format (With inputs wrapper) ===")
response = h({
"inputs": {
"model": "Ephraimmm/PIDGIN_gemma-3",
"messages": [
{"role": "user", "content": "How far?"}
],
"temperature": 0.7,
"max_tokens": 64
}
})
print(response)
print()
# Test 3: Multi-turn conversation
print("=== Test 3: Multi-turn Conversation ===")
response = h({
"model": "Ephraimmm/PIDGIN_gemma-3",
"messages": [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Explain why the sky is blue"},
{"role": "assistant", "content": "The sky dey blue because..."},
{"role": "user", "content": "Tell me more"}
],
"max_tokens": 100
})
print(response)