Download code/duovlm.py from Duoia/duovlm-40m-v1: direct link, hf CLI and curl.
- Browser
- Download file 9.72 kB
-
https://huggingface.co/Duoia/duovlm-40m-v1/resolve/main/code/duovlm.py
- Command line
-
hf download hf://Duoia/duovlm-40m-v1/code/duovlm.py
-
curl -L -o duovlm.py https://huggingface.co/Duoia/duovlm-40m-v1/resolve/main/code/duovlm.py
9.72 kB
| #!/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 | |
| 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")} | |