Diba-Vision / handler.py
DibaAi's picture
Release
ffd2e20
Raw History Blame Contribute Delete
4.28 kB
"""
Diba-Vision Inference Endpoint handler. Accepts text and image(s), replies in the user's language.
Request body (any of):
{"inputs": "توضیح بده"} # text only
{"inputs": "این عکس چیست؟", "image": "data:image/png;base64,..."} # text + one image (data URI, bare base64, or URL)
{"inputs": [ {"role":"user","content":[{"type":"image","image":"..."},{"type":"text","text":"..."}]} ]} # full chat
{"parameters": {"max_new_tokens": 512, "temperature": 0}}
"""
import base64
import io
import os
import torch
from PIL import Image
from transformers import AutoModelForImageTextToText, AutoProcessor
SYSTEM = {
"fa": "تو «دیبا» هستی، دستیار هوش مصنوعی دیباچین. به همان زبانی که کاربر نوشته پاسخ بده؛ فارسی نوشتاری، روشن و دقیق. اگر تصویری هست، دقیق بر اساس همان تصویر جواب بده و چیزی نساز.",
"en": "You are Diba, an AI assistant by Dibachain. Reply in the language the user wrote in, clearly and accurately. When an image is present, describe it faithfully and do not invent details.",
}
def _lang(text):
fa = sum(1 for c in text if "؀" <= c <= "ۿ")
en = sum(1 for c in text if c.isascii() and c.isalpha())
return "fa" if fa >= en else "en"
def _image(spec):
if isinstance(spec, dict):
spec = spec.get("image") or spec.get("url") or spec.get("image_url") or ""
if not isinstance(spec, str) or not spec:
return None
if spec.startswith("data:"):
spec = spec.split(",", 1)[1]
if spec.startswith(("http://", "https://")):
return spec
try:
return Image.open(io.BytesIO(base64.b64decode(spec))).convert("RGB")
except Exception:
return None
class EndpointHandler:
def __init__(self, path=""):
tok = os.environ.get("HF_TOKEN")
self.processor = AutoProcessor.from_pretrained(path, trust_remote_code=True, token=tok)
self.model = AutoModelForImageTextToText.from_pretrained(
path, trust_remote_code=True, dtype=torch.bfloat16, token=tok, device_map="auto").eval()
def _messages(self, data):
inp = data.get("inputs", data.get("text", ""))
if isinstance(inp, list): # already chat messages
msgs = inp
last = " ".join(p.get("text", "") for m in msgs if isinstance(m.get("content"), list)
for p in m["content"] if p.get("type") == "text")
else:
parts = []
img = _image(data.get("image") or data.get("image_url"))
if img is not None:
parts.append({"type": "image", "image": img})
parts.append({"type": "text", "text": str(inp)})
msgs = [{"role": "user", "content": parts}]
last = str(inp)
# resolve any string image specs inside provided messages
for m in msgs:
if isinstance(m.get("content"), list):
for p in m["content"]:
if p.get("type") == "image" and isinstance(p.get("image"), str):
im = _image(p["image"])
if im is not None:
p["image"] = im
sysmsg = {"role": "system", "content": [{"type": "text", "text": SYSTEM[_lang(last)]}]}
return [sysmsg] + msgs
def __call__(self, data):
params = data.get("parameters") or {}
messages = self._messages(data)
inputs = self.processor.apply_chat_template(
messages, add_generation_prompt=True, tokenize=True, enable_thinking=False,
return_dict=True, return_tensors="pt").to(self.model.device)
temp = float(params.get("temperature", 0) or 0)
with torch.no_grad():
out = self.model.generate(
**inputs, max_new_tokens=int(params.get("max_new_tokens", 512)),
do_sample=temp > 0, temperature=max(temp, 1e-5), top_p=float(params.get("top_p", 0.9)),
repetition_penalty=1.05)
text = self.processor.batch_decode(out[:, inputs["input_ids"].shape[1]:], skip_special_tokens=True)[0].strip()
return [{"generated_text": text}]