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")}
|