File size: 4,276 Bytes
ffd2e20
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
"""
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}]