""" 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}]