【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, bias=False),其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做 flatten,lm_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(hidden=8, vocab=16) hidden = torch.randn(2, 4, 8) # [B, T, hidden] # 正常情况 out = head.logits(hidden) print("正常 2-D weight,logits 形状:", 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, bias=False) 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, vocab=16, hidden=8) 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( f"lm_head.weight 必须是 2 维 [vocab, hidden]," f"当前是 {weight.dim()} 维 {tuple(weight.shape)}" ) if tuple(weight.shape) != (self.vocab, self.hidden): raise ValueError( f"lm_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, bias=False) def generate(self, hidden): # generate 入口先 sanitize return torch.matmul(hidden, self.lm_head.weight.T) def main(): spec = HeadSpec(vocab=16, hidden=8) 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, bias=False) 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 前还原。