File size: 9,723 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
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
#!/usr/bin/env python3
"""DuoVLM-40M 模型包装:CLIP(冻结) + 连接器 + 自训 GPT。

锁定设计(PLAN_VLM_40M.md §3.1):

    <|bos|> [IMG]×196 <|user|> {question} <|eot|> <|assistant|> {answer} <|eos|>
            ↑ 196 个图像占位 token 的 embedding 被连接器输出直接覆写;只有答案段算 loss

为什么不用 litgpt 的 GPT.forward:它只接受 token id(§8 集成点 2),
本文件手写 wte → 图像注入 → blocks → ln_f → lm_head 这条路径。

踩过的坑(都写在 §8):
  1. 权重绑定:litgpt 只在 pretrain 路径 tie,这里必须手工绑,
     否则参数量打印成 42.2M 超过 40M 预算(数字误导)。
  2. 位移约定:litgpt 的 chunked_cross_entropy 不做位移,
     调用方要自己错开一位:CE(logits[:, :-1], labels[:, 1:])。
"""
import torch
import torch.nn as nn

from litgpt.config import Config
from litgpt.model import GPT
from litgpt.utils import chunked_cross_entropy

# tokenizer/tokenizer.json 实测的 8 个特殊 token
BOS, EOS, EOT, PAD, USER, ASSISTANT, SYSTEM, IMAGE = 0, 1, 2, 3, 4, 5, 6, 7
N_IMG = 196
IMG_START = 1  # BOS 之后立刻是 196 个图像位(固定,故可用切片拼接注入)


class DuoVLM(nn.Module):
    def __init__(self, cfg: Config, c_in: int = 768, c_hid: int = 1024) -> None:
        super().__init__()
        self.cfg = cfg
        self.llm = GPT(cfg)
        # 连接器:768 → 1024 → 512(GELU),1,312,256 参数
        self.connector = nn.Sequential(
            nn.Linear(c_in, c_hid),
            nn.GELU(),
            nn.Linear(c_hid, cfg.n_embd),
        )
        # 权重绑定(必须,见文件头坑 1)
        self.llm.lm_head.weight = self.llm.transformer.wte.weight

    # ---- 参数账 ----
    def param_report(self) -> str:
        llm = sum(p.numel() for p in self.llm.parameters())
        con = sum(p.numel() for p in self.connector.parameters())
        return (
            f"LLM {llm/1e6:.3f}M + connector {con/1e6:.3f}M = "
            f"{(llm+con)/1e6:.3f}M 可训练参数"
        )

    def set_trainable(self, llm: bool, connector: bool = True) -> None:
        for p in self.llm.parameters():
            p.requires_grad_(llm)
        for p in self.connector.parameters():
            p.requires_grad_(connector)

    # ---- 前向 ----
    def forward(
        self,
        ids: torch.Tensor,          # (B, T) int64
        feats: torch.Tensor | None = None,   # (B, 196, 768) CLIP 特征(fp16/bf16/fp32)
        labels: torch.Tensor | None = None,  # (B, T) int64,非答案位 = -100
    ):
        x = self.llm.transformer.wte(ids)
        if self.cfg.scale_embeddings:
            x = x * torch.tensor(self.cfg.n_embd**0.5, dtype=x.dtype)

        if feats is not None:
            v = self.connector(feats.to(x.dtype))
            # 模板固定:图像位恒在 [IMG_START, IMG_START+196),用切片拼接注入
            # (等价于把 196 个占位 embedding 覆写成连接器输出)
            assert ids.size(1) >= IMG_START + N_IMG, "序列比 196 个图像位还短"
            x = torch.cat([x[:, :IMG_START], v, x[:, IMG_START + N_IMG:]], dim=1)

        cos = self.llm.cos[: ids.size(1)].unsqueeze(0)
        sin = self.llm.sin[: ids.size(1)].unsqueeze(0)
        for block_idx, block in enumerate(self.llm.transformer.h):
            if self.cfg.rope_indices is not None:
                x = block(
                    x,
                    cos[..., self.cfg.rope_indices[block_idx]],
                    sin[..., self.cfg.rope_indices[block_idx]],
                    None,
                    None,
                    None,
                )
            else:
                x = block(x, cos, sin, None, None, None)
        x = self.llm.transformer.ln_f(x)
        logits = self.llm.lm_head(x)

        loss = None
        if labels is not None:
            # 位移:logits[t] 预测 label[t+1](见文件头坑 2)
            loss = chunked_cross_entropy(logits[:, :-1, :], labels[:, 1:], chunk_size=128)
        return logits, loss

    @torch.no_grad()
    def generate(
        self,
        ids: torch.Tensor,               # (B, T_prompt)
        feats: torch.Tensor,             # (B, 196, 768)
        max_new_tokens: int = 16,
        no_repeat_ngram: int = 3,
        eos_id: int = EOS,
    ) -> list[list[int]]:
        """贪心解码 + n-gram 复读抑制(全量重算,不建 KV cache)。

        实测结论(STAGE1_ROADMAP §4.5):greedy 无约束时模型 100% 复读、
        从不吐 <|eos|>;加 no_repeat_ngram=3 后退化 0/5。所以这里是必需项。
        """
        self.eval()
        out = [[] for _ in range(ids.size(0))]
        done = [False] * ids.size(0)
        cur = ids
        for _ in range(max_new_tokens):
            logits = self(cur, feats)[0][:, -1, :].float()  # (B, V)
            for b in range(cur.size(0)):
                if done[b]:
                    continue
                hist = out[b]
                if no_repeat_ngram > 0 and len(hist) >= no_repeat_ngram - 1:
                    n = no_repeat_ngram
                    prefix = tuple(hist[-(n - 1):]) if n > 1 else ()
                    banned = set()
                    for i in range(len(hist) - n + 1):
                        if tuple(hist[i:i + n - 1]) == prefix and i + n - 1 < len(hist):
                            banned.add(hist[i + n - 1])
                    if banned:
                        logits[b, list(banned)] = -float("inf")
                nxt = int(torch.argmax(logits[b]).item())
                hist.append(nxt)
                if nxt == eos_id:
                    done[b] = True
            cur = torch.cat(
                [cur, torch.tensor([[out[b][-1] if out[b] else 0] for b in range(cur.size(0))],
                                   device=cur.device, dtype=cur.dtype)],
                dim=1,
            )
            if all(done):
                break
        return out


def build_sequence(ids_tok: list[int], question: str, answer: str, tokenizer, max_len: int = 512):
    """按 §3.1 模板拼一条序列,返回 (ids, labels, n_truncated_answer_tokens)。

    模板: BOS + [IMAGE]*196 + USER + q + EOT + ASSISTANT + a + EOS
    labels: 只有答案段 + EOS 位非 -100(其它含 196 图像位全部 -100)
    """
    q = tokenizer.encode(question).tolist() if question else []
    a = tokenizer.encode(answer).tolist() if answer else []

    # 固定部分:BOS(1) + 196 + USER(1) + EOT(1) + ASSISTANT(1) = 200
    fixed = 1 + N_IMG + 3
    budget = max_len - fixed - 1  # 留 1 给 EOS
    q = q[: max(0, budget - 1)]   # 至少给答案留 1 个 token
    budget_a = budget - len(q)
    n_cut = max(0, len(a) - budget_a)
    a = a[:budget_a]

    ids = [BOS] + [IMAGE] * N_IMG + [USER] + q + [EOT, ASSISTANT] + a + [EOS]
    labels = [-100] * len(ids)
    ans_start = 1 + N_IMG + 3 + len(q)          # 第一个答案 token 的下标
    for i in range(ans_start, len(ids)):
        labels[i] = ids[i]
    assert ids[IMG_START:IMG_START + N_IMG] == [IMAGE] * N_IMG
    return ids, labels, n_cut


# ---- VQA 答案归一化(§5.1 规则 1:子集只有单条答案,只能用严格匹配)----
import re  # noqa: E402

_ART = re.compile(r"\b(a|an|the)\b")
_PUNC = re.compile(r"[^\w\s]")


def norm_answer(s: str) -> str:
    s = s.lower().strip()
    s = _PUNC.sub(" ", s)
    s = _ART.sub(" ", s)
    return " ".join(s.split())


# ---- 权重存取(自定格式,存 llm / connector 两份,便于 Stage 3 从 Stage 2 续)----
def save_duovlm(path, model: "DuoVLM", step: int, extra: dict | None = None) -> None:
    """存 llm / connector 两份权重。

    ⚠️ 坑:self.llm.state_dict() 的键本来就没有 'llm.' 前缀(它是子模块自己的 state_dict),
    早期版本画蛇添足地剥了 4 个字符,把 'transformer...' 存成 'sformer...'、
    'lm_head.weight' 存成 'ead.weight'——数值没错但键名全废。load_duovlm 里做了兼容还原。
    """
    from pathlib import Path

    p = Path(path)
    p.parent.mkdir(parents=True, exist_ok=True)
    torch.save(
        {
            "llm": dict(model.llm.state_dict()),
            "connector": model.connector.state_dict(),
            "step": step,
            "extra": extra or {},
        },
        p,
    )


def _repair_llm_keys(sd: dict, model: "DuoVLM") -> dict:
    """还原被截断 4 个字符的历史键名;对不上就断言失败,绝不静默错配。"""
    good = list(model.llm.state_dict().keys())
    out = {}
    for k, v in sd.items():
        if k in good:
            out[k] = v
            continue
        cands = [mk for mk in good if len(mk) - len(k) == 4 and mk.endswith(k)]
        assert len(cands) == 1, f"无法还原键 {k!r}(候选 {cands})"
        out[cands[0]] = v
    return out


def load_duovlm(path, model: "DuoVLM", load_llm: bool = True, load_connector: bool = True) -> dict:
    d = torch.load(path, map_location="cpu")
    if load_llm:
        llm_sd = _repair_llm_keys(d["llm"], model)
        msd = model.llm.state_dict()
        for k, v in llm_sd.items():
            assert v.shape == msd[k].shape, f"形状不匹配 {k}: {tuple(v.shape)} vs {tuple(msd[k].shape)}"
        miss, unexp = model.llm.load_state_dict(llm_sd, strict=False)
        assert not unexp, f"llm 权重有多余键: {unexp}"
    if load_connector and "connector" in d:
        model.connector.load_state_dict(d["connector"])
    # 绑定可能被 load_state_dict 破坏(tie 过的两份权重会被分别赋值)
    model.llm.lm_head.weight = model.llm.transformer.wte.weight
    return d.get("extra", {}) | {"step": d.get("step")}