File size: 7,707 Bytes
4afe981
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
#!/usr/bin/env python3
"""DuoVLM-40M 推理封装(包内自包含:给一张图 + 一个问题 → 回答)

路径全部可用环境变量覆盖,默认相对包根:
  DUOVLM_WEIGHTS    权重 .pth       默认 <包根>/weights/duovlm-s3-final.pth
  DUOVLM_CFG        模型配置        默认 <权重同目录>/model_config.yaml
  DUOVLM_TOKENIZER  词表目录        默认 <包根>/tokenizer
  DUOVLM_CLIP       视觉塔 HF id    默认 openai/clip-vit-base-patch16
  DUOVLM_CLIP_DIR   视觉塔本地目录  设了就优先用它(离线场景)
  DUOVLM_DEVICE     cuda / cpu      默认自动

依赖:torch、transformers、litgpt==0.5.13(见 requirements.txt)。
"""
from __future__ import annotations

import os
import time
from contextlib import nullcontext
from pathlib import Path

import numpy as np
import torch
from PIL import Image

HERE = Path(__file__).resolve().parent
PKG = HERE.parent
DEFAULT_CLIP = "openai/clip-vit-base-patch16"

from decoding import decode_safe                                              # noqa: E402
from duovlm import (ASSISTANT, BOS, EOS, EOT, IMAGE, N_IMG, DuoVLM,           # noqa: E402
                    load_duovlm)

DESC_Q = "Render a clear and concise summary of the photo."


def resolve_paths(weights: str | None = None) -> dict:
    w = Path(weights or os.environ.get("DUOVLM_WEIGHTS") or PKG / "weights" / "duovlm-s3-final.pth")
    return {
        "weights": w,
        "cfg": Path(os.environ.get("DUOVLM_CFG") or (w.parent / "model_config.yaml")),
        "tokenizer": Path(os.environ.get("DUOVLM_TOKENIZER") or PKG / "tokenizer"),
    }


class DuoVLMInfer:
    """一次性载入(约 5~15s),之后每次问答 0.1~1.8s(取决于是否要现跑视觉塔)。"""

    def __init__(self, weights: str | None = None, device: str | None = None,
                 clip: str | None = None, verbose: bool = True):
        from litgpt.config import Config
        from litgpt.tokenizer import Tokenizer

        self.paths = resolve_paths(weights)
        self.dev = device or os.environ.get("DUOVLM_DEVICE") or ("cuda" if torch.cuda.is_available() else "cpu")
        self.dtype = torch.bfloat16 if self.dev == "cuda" else torch.float32
        self.tok = Tokenizer(self.paths["tokenizer"])
        try:
            cfg = Config.from_file(str(self.paths["cfg"]))
        except Exception:
            import yaml
            cfg = Config(**yaml.safe_load(self.paths["cfg"].read_text()))
        self.model = DuoVLM(cfg)
        info = load_duovlm(self.paths["weights"], self.model)
        self.model = self.model.to(self.dev, dtype=self.dtype).eval()
        self.step = info.get("step")
        self.extra = {k: (float(v) if hasattr(v, "item") else v)
                      for k, v in (info.get("extra") or {}).items()}
        if not self.extra:      # load_duovlm 不一定回传 extra,这里从权重直接补读(只读元数据)
            try:
                _d = torch.load(self.paths["weights"], map_location="cpu")
                self.step = self.step or _d.get("step")
                self.extra = {k: (float(v) if hasattr(v, "item") else v)
                              for k, v in (_d.get("extra") or {}).items()}
            except Exception:
                pass
        self.n_params = sum(x.numel() for x in self.model.parameters())
        self.clip_source = clip or os.environ.get("DUOVLM_CLIP_DIR") or os.environ.get("DUOVLM_CLIP") or DEFAULT_CLIP
        self._clip = None
        if verbose:
            print(f"[duovlm] 权重 {self.paths['weights'].name}(step {self.step})"
                  f"  参数 {self.n_params:,}  设备 {self.dev}  dtype {self.dtype}")
            print(f"[duovlm] 视觉塔 {self.clip_source}(首次用到时下载/载入)", flush=True)

    def param_report(self) -> str:
        llm = sum(p.numel() for p in self.model.llm.parameters())
        con = sum(p.numel() for p in self.model.connector.parameters())
        return (f"LLM {llm:,} + connector {con:,} = {llm + con:,} 可训练参数")

    # ---------- 视觉塔 ----------
    def _load_clip(self):
        if self._clip is None:
            from transformers import CLIPImageProcessor, CLIPVisionModel

            self.proc = CLIPImageProcessor.from_pretrained(self.clip_source)
            self._clip = CLIPVisionModel.from_pretrained(self.clip_source,
                                                         dtype=self.dtype).to(self.dev).eval()
        return self._clip

    @torch.no_grad()
    def embed(self, image) -> np.ndarray:
        """返回 (196, 768) fp16 特征。image 可以是路径或 PIL.Image。"""
        clip = self._load_clip()
        im = Image.open(image) if isinstance(image, (str, Path)) else image
        px = self.proc(im.convert("RGB"), return_tensors="pt")["pixel_values"]
        ctx = torch.autocast("cuda", dtype=torch.bfloat16) if self.dev == "cuda" else nullcontext()
        with torch.no_grad(), ctx:
            h = clip(pixel_values=px.to(self.dev, dtype=self.dtype),
                     output_hidden_states=True).hidden_states[-2]
            h = clip.vision_model.post_layernorm(h)[:, 1:, :]
        return h[0].float().cpu().numpy().astype(np.float16)

    # ---------- 生成 ----------
    @torch.no_grad()
    def ask(self, image=None, question: str = "", blind: bool = False, no_image: bool = False,
            max_new: int = 12, ngram: int = 3, rep_penalty: float = 1.0,
            temperature: float = 0.0, top_k: int = 0, loop_break: bool = True) -> dict:
        """看图问答。blind=True 表示保留模板但图像位换零(不看图对照);
        no_image=True 表示连 196 个图像位都去掉(纯文本,分布外)。"""
        q = (question or DESC_Q).replace("<image>", " ").strip()
        qids = self.tok.encode(q).tolist()
        feat, src = None, "no-image"
        if no_image:
            ids = [BOS] + qids[: 512 - 4] + [EOT, ASSISTANT]
        else:
            ids = [BOS] + [IMAGE] * N_IMG + [4] + qids[: 512 - (N_IMG + 4)] + [EOT, ASSISTANT]
            t_emb = time.time()
            feat = self.embed(image)
            src = "blind(zeros)" if blind else "image"
            if blind:
                feat = np.zeros_like(feat)
            emb_ms = (time.time() - t_emb) * 1000
        t0 = time.time()
        ft = None if feat is None else torch.from_numpy(feat[None]).to(self.dev)
        out, st = decode_safe(self.model, torch.tensor([ids], dtype=torch.long).to(self.dev), ft,
                              max_new=max_new, ngram=ngram, rep_penalty=rep_penalty,
                              temperature=temperature, top_k=top_k, loop_break=loop_break)
        toks = out[0]
        ans = " ".join(self.tok.decode(torch.tensor(toks)).split()) if toks else ""
        return {"answer": ans, "tokens": len(toks), "ms": round((time.time() - t0) * 1000),
                "emb_ms": round(locals().get("emb_ms", 0)), "source": src,
                "stats": {k: v for k, v in st.items() if v}}

    @torch.no_grad()
    def continue_text(self, seed: str, max_new: int = 40, **kw) -> dict:
        """纯文本续写(无图像位,Stage 1 原生格式)。"""
        ids = ([BOS] + self.tok.encode(seed).tolist())[: 512 - max_new - 1]
        t0 = time.time()
        out, st = decode_safe(self.model, torch.tensor([ids], dtype=torch.long).to(self.dev),
                              None, max_new=max_new, **kw)
        toks = out[0]
        return {"answer": " ".join(self.tok.decode(torch.tensor(toks)).split()) if toks else "",
                "tokens": len(toks), "ms": round((time.time() - t0) * 1000), "source": "text",
                "stats": {k: v for k, v in st.items() if v}}