ier_qwen3MergedModel / handler.py
Infraizoo's picture
Create handler.py
0cd6415 verified
Raw
History Blame Contribute Delete
1.87 kB
from transformers import AutoTokenizer, AutoModelForVision2Seq
import torch, os
from jinja2 import Template
from typing import Any, Dict, List
class EndpointHandler:
def __init__(self, model_dir: str = "", **kwargs: Any):
# Load tokenizer
self.tokenizer = AutoTokenizer.from_pretrained(model_dir, use_fast=True, trust_remote_code=True)
# Load model with trust_remote_code to handle custom Qwen classes
self.model = AutoModelForVision2Seq.from_pretrained(
model_dir,
torch_dtype=torch.float16 if torch.cuda.is_available() else torch.float32,
device_map="auto",
trust_remote_code=True,
)
self.model.eval()
# Load chat template
template_path = os.path.join(model_dir, "chat_template.jinja")
with open(template_path, "r", encoding="utf-8") as f:
self.template = Template(f.read())
def _render_prompt(self, messages: List[Dict[str, Any]], tools=None):
return self.template.render(
messages=messages,
tools=tools or [],
add_generation_prompt=True,
add_vision_id=False,
)
def __call__(self, data: Dict[str, Any]) -> Dict[str, Any]:
messages = data.get("messages", [])
tools = data.get("tools", None)
prompt = self._render_prompt(messages, tools)
inputs = self.tokenizer(prompt, return_tensors="pt").to(self.model.device)
gen_kwargs = {
"max_new_tokens": data.get("max_new_tokens", 256),
"temperature": data.get("temperature", 0.7),
"top_p": data.get("top_p", 0.9),
}
with torch.no_grad():
output = self.model.generate(**inputs, **gen_kwargs)
text = self.tokenizer.decode(output[0], skip_special_tokens=True)
return {"generated_text": text}