duovlm-40m-v1 / code /duovlm.py
Duoia's picture
DuoVLM-40M v1: from-scratch 40M vision-language model (frozen CLIP + MiniPile-pretrained LM)
4afe981 verified
Raw History Blame Contribute Delete
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
@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")}