#!/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")}