多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

【Bug已解决】Error generating example: ‘weight‘ must be 2-D in model.generate() 解决方案

【Bug已解决】Error generating example: ‘weight‘ must be 2-D in model.generate() 解决方案 【Bug已解决】Error generating example weight must be 2-D in model.generate() 解决方案一、现象长什么样用model.generate(...)做文本生成时模型加载、forward 都正常一进生成就炸RuntimeError: weight must be 2-D或者更完整一点RuntimeError: weight must be 2-D, but got weight of shape [torch.Size([25600])]有时还伴随在transformers的CausalLM里generate调用lm_head计算下一个 token 的 logits 时失败而正常的model(input_ids).logits却没问题——这种forward 正常、generate 报错的差异最让人困惑。现象的本质是model.generate()内部需要反复调用lm_head输出投影层把隐藏状态映射成词表 logits而lm_head的权重weight在这个时刻不是 2 维[vocab, hidden]变成了 1 维或被错误 reshape 了于是F.linear(hidden, weight)直接拒绝。二、背景lm_head本质是一个nn.Linear(hidden, vocab, biasFalse)其weight形状应为[vocab, hidden]2 维。F.linear(x, w)要求w是 2 维。生成时transformers的CausalLM在prepare_inputs_for_generation之后用lm_head(hidden_states)算 logits。什么情况下weight会变 1 维量化/合并merge的副作用用 bitsandbytes / GPTQ / AWQ 量化或把 LoRA 合并进基座后某些代码为了省显存把lm_head.weight做了.view(-1)/.flatten()或在state_dict往返时丢了形状信息。generate 时又没恢复 2 维。tie 权重处理不当lm_head.weight embed_tokens.weight共享而embed_tokens是[vocab, hidden]没问题但若有人对embed_tokens做了weight.flatten().view(...)之类的优化共享的lm_head.weight也就跟着变成 1 维。FSDP2 / TP 分片后的视图错误分片把weight切成 DTensor 的 local 切片若.to_local()后形状被错误地squeeze/flatten恢复 2 维的视图没建好。自定义 generate 逻辑误 reshape用户在compute_logits里手写了weight.view(-1)之类。下面用可运行代码复现lm_head.weight 变 1 维导致F.linear报 weight must be 2-D。三、根因根因一句话lm_head.weight在进入model.generate()时被错误地弄成了非 2 维通常是 1 维 flattened而F.linear要求权重 2 维于是 generate 报weight must be 2-D。三个具体失配量化/合并把 weight flatten 成 1 维为了紧凑存储合并后.view(-1)generate 前未恢复[vocab, hidden]。tie 权重共享被连带 reshape对embed_tokens做 flattenlm_head.weight因共享变成 1 维。分片 local 视图恢复缺失FSDP2/TP 切分后.to_local()形状错乱没重建 2 维视图。四、最小可运行复现用一段纯torch模拟lm_head的F.linear调用先正常 2 维、再把weight错误 flatten 成 1 维复现报错import torch import torch.nn as nn import torch.nn.functional as F class TinyLMHead(nn.Module): def __init__(self, hidden, vocab): super().__init__() self.weight nn.Parameter(torch.randn(vocab, hidden)) # [vocab, hidden] 2-D def logits(self, hidden): return F.linear(hidden, self.weight) # 要求 weight 2-D def main(): head TinyLMHead(hidden8, vocab16) hidden torch.randn(2, 4, 8) # [B, T, hidden] # 正常情况 out head.logits(hidden) print(正常 2-D weightlogits 形状:, tuple(out.shape)) # 错误情况weight 被 flatten 成 1-D模拟合并/量化副作用 bad head.weight.data.flatten().clone() head.weight nn.Parameter(bad) # [vocab*hidden] 1-D try: head.logits(hidden) except RuntimeError as e: print(复现到报错:, e) if __name__ __main__: main()运行会先打印正常形状再打印复现到报错: weight must be 2-D, but got weight of shape ...[128]——正是 generate 时报错的本质。五、解决方案第一层最小直接修复最立竿见影的修复确保lm_head.weight在 generate 之前恢复成[vocab, hidden]的 2 维视图。如果是被 flatten 了用.view(vocab, hidden)恢复如果是因为 tie确保embed_tokens不被 flatten。import torch import torch.nn as nn def ensure_lm_head_2d(model, vocab, hidden): 修复把 lm_head.weight 强制恢复成 2 维 [vocab, hidden]。 w model.lm_head.weight if w.dim() ! 2: # 展平后按 vocab x hidden 重排优先用 .view共享视图省显存 model.lm_head.weight nn.Parameter(w.reshape(vocab, hidden)) return model class TinyLM(nn.Module): def __init__(self, hidden, vocab): super().__init__() self.hidden hidden self.vocab vocab self.embed nn.Parameter(torch.randn(vocab, hidden)) self.lm_head nn.Linear(hidden, vocab, biasFalse) self.lm_head.weight self.embed # tie def generate_step(self, hidden): return torch.matmul(hidden, self.lm_head.weight.T) # [B,T,vocab] def main(): model TinyLM(8, 16) # 假设合并/量化把 embed 错误 flatten 了连带 lm_head 也变 1 维 flat model.embed.data.flatten().clone() model.embed nn.Parameter(flat) model.lm_head.weight model.embed ensure_lm_head_2d(model, vocab16, hidden8) out model.generate_step(torch.randn(2, 4, 8)) print(修复后生成 logits 形状:, tuple(out.shape)) if __name__ __main__: main()第一层修复直接在 generate 前把weight恢复 2 维报错消失。六、解决方案第二层结构性改进把lm_head.weight 必须 2 维收口成一个HeadSanitizer在模型构建完成、以及在generate调用入口处强制校验并修复形状避免任何 flatten 漏网。import torch import torch.nn as nn from dataclasses import dataclass dataclass class HeadSpec: vocab: int hidden: int def assert_2d(self, weight: torch.Tensor): if weight.dim() ! 2: raise ValueError( flm_head.weight 必须是 2 维 [vocab, hidden] f当前是 {weight.dim()} 维 {tuple(weight.shape)} ) if tuple(weight.shape) ! (self.vocab, self.hidden): raise ValueError( flm_head.weight 形状应为 {(self.vocab, self.hidden)} f实际 {tuple(weight.shape)} ) def sanitize(self, model: nn.Module) - nn.Module: w model.lm_head.weight if w.dim() ! 2 or tuple(w.shape) ! (self.vocab, self.hidden): # 自动恢复从展平/错误形状重建 2 维视图 model.lm_head.weight nn.Parameter(w.reshape(self.vocab, self.hidden)) else: self.assert_2d(model.lm_head.weight) return model class TinyLM(nn.Module): def __init__(self, hidden, vocab): super().__init__() self.lm_head nn.Linear(hidden, vocab, biasFalse) def generate(self, hidden): # generate 入口先 sanitize return torch.matmul(hidden, self.lm_head.weight.T) def main(): spec HeadSpec(vocab16, hidden8) model TinyLM(8, 16) # 模拟被 flatten 的 weight model.lm_head.weight nn.Parameter(model.lm_head.weight.data.flatten().clone()) spec.sanitize(model) out model.generate(torch.randn(2, 4, 8)) print(结构层修复后 generate 正常形状:, tuple(out.shape)) if __name__ __main__: main()第二层的关键是HeadSpec把2 维约束 自动恢复固化在 generate 之前的必经路径任何 reshape 错误都会被拦截或自动修好。七、解决方案第三层断言 / CI 守护加 pytest 守护(1) 正常 2 维 weight 可通过F.linear(2) 1 维 weight 必须被 sanitizer 恢复(3) generate 入口拒绝非 2 维 weight。用纯 torch 模拟import torch import torch.nn as nn import torch.nn.functional as F import pytest def logits(weight, hidden): return F.linear(hidden, weight) def test_2d_weight_passes(): w torch.randn(16, 8) out logits(w, torch.randn(2, 4, 8)) assert out.shape (2, 4, 16) def test_1d_weight_raises(): w torch.randn(16 * 8) # 1-D with pytest.raises(RuntimeError): logits(w, torch.randn(2, 4, 8)) def test_sanitizer_restores_2d(): class M(nn.Module): def __init__(self): super().__init__() self.lm_head nn.Linear(8, 16, biasFalse) m M() # 破坏成 1-D m.lm_head.weight nn.Parameter(m.lm_head.weight.data.flatten().clone()) assert m.lm_head.weight.dim() ! 2 # 恢复 m.lm_head.weight nn.Parameter(m.lm_head.weight.reshape(16, 8)) assert m.lm_head.weight.dim() 2 out logits(m.lm_head.weight, torch.randn(2, 4, 8)) assert out.shape (2, 4, 16) if __name__ __main__: pytest.main([__file__, -q])CI 里test_1d_weight_raises验证1 维必报错这个不变量test_sanitizer_restores_2d验证自动恢复有效从根上防住 generate 时的 2-D 报错回归。八、排查清单model.generate()报weight must be 2-D时按此顺序查先确认是不是 generate 专属若model(input_ids).logits正常但generate报错基本锁定lm_head.weight形状问题generate 反复调 lm_head。打印model.lm_head.weight.shape确认是不是[vocab, hidden]的 2 维。不是就找到了根。回想是否做过量化/合并LoRA 合并、GPTQ/AWQ/bnb 量化后是否对lm_head.weight或embed_tokens做过.view(-1)/flatten。有的话在 generate 前恢复 2 维。检查 tie若lm_head.weight is embed_tokens.weight检查embed_tokens是否被 reshape 连带影响。检查 FSDP2/TP 分片恢复分片后.to_local()的 local 形状是否正确重建了 2 维视图。生成前加断言在generate入口加assert model.lm_head.weight.dim() 2把隐患变成显式报错。优先用.view而非新建恢复 2 维时尽量用共享视图.view(vocab, hidden)避免额外显存与拷贝。九、小结model.generate()报weight must be 2-D根因不是生成逻辑坏了而是lm_head.weight在被反复调用算 logits 时已不是合法的 2 维[vocab, hidden]——通常是量化/LoRA 合并时把权重 flatten 成 1 维、或 tie 权重被连带 reshape、或分片 local 视图恢复缺失。F.linear明确要求权重 2 维于是 generate 在第一次算 logits 时就炸而普通 forward 可能因走的是另一条路径而正常造成forward 行、generate 不行的迷惑。修复三层第一层在 generate 前用.view(vocab, hidden)把 weight 恢复 2 维第二层用HeadSpec把2 维约束 自动恢复固化在 generate 必经路径第三层用 pytest 断言1 维必报错、sanitizer 能恢复。记住lm_head.weight永远是[vocab, hidden]的 2 维任何 flatten 都必须在 generate 前还原。
返回列表