【Bug已解决】Training collapses when token embeddings are not padded to a multiple of 64 解决方案
一、现象长什么样
在用一个自定义词表(vocab_size 不是整数)做训练时,只要vocab_size不是64 的倍数,训练就会在几十步内崩塌:loss 突然变成nan/inf,或 loss 先正常下降随后剧烈 spike 到极大值,模型参数迅速变废。
典型日志:
step 30: loss=2.1 (正常) step 35: loss=nan # 或 step 40: loss=87.3 (spike) -> 之后全 nan而把 vocab_size 向上凑成 64 的倍数(比如 50000 → 50016)后,同样的配置、同样的数据,训练稳稳收敛。
现象特征:
- 只在 vocab_size 非 64 倍数时崩,是确定性复现(不是随机);
- 崩得快(几十步内),不是缓慢退化;
- 崩的时机在"第一次大规模 embedding/lm_head 矩阵乘之后",指向词表维度对齐问题。
这是典型的"硬件/内核要求张量维度按 64 对齐,未对齐导致数值/内存错乱"。
二、背景
现代 GPU 的 Tensor Core 要求矩阵维度按8/16/64对齐才能达到峰值效率;更关键的是,很多融合内核和 padding 逻辑假设词表维度是 64 的倍数:
- Tensor Core 对齐:
[B, vocab]这类大矩阵乘,vocab 非 64 倍数时,内核内部会做 padding,若上层没配套处理 padding 行,padding 区域的 garbage 会混入结果。 - lm_head 与 embedding 共享:
nn.Embedding(vocab, d)和Linear(d, vocab)在 vocab 非对齐时,权重张量的实际分配可能被框架/内核 pad 到 64 倍数,但前向计算没 mask 掉 padding 行对应的 logits,于是 padding token 的 logits 参与 softmax,概率被稀释、且 padding 行的未初始化值污染分布。 - 序列并行 / 专家并行:在 TP/EP 下,词表被切分到多卡,切分要求能被
world_size和 64 同时整除,否则边界行越界或错位,产生 NaN。 - loss 计算未屏蔽 padding:即使前向 pad 了 vocab,CrossEntropy 的
ignore_index只屏蔽了序列位置的 padding(token id=-100),没屏蔽词表维度的 padding 行,于是这 64 个(或更少)padding 行成了"幽灵类别",它们的 logits 参与 softmax,数值失真、训练崩。
三、根因
根因一句话:词表维度vocab_size不是 64 的倍数时,embedding/lm_head 的矩阵乘内核做了内部 padding,但上层既没有把 vocab pad 到 64 倍数并对 padding 行做 logits 屏蔽,也没有在 CrossEntropy 里忽略这些 padding 类别,导致未初始化的 padding 行 logits 污染 softmax 分布,训练在几十步内因数值错乱而崩塌。
具体:
- 未对齐:vocab 非 64 倍数,Tensor Core / 融合内核内部 pad 出若干 padding 行;
- 未屏蔽:padding 行对应的 logits 没被置为
-inf或忽略,参与 softmax; - 污染分布:padding 行的 garbage 值让 softmax 概率失真,梯度方向错,loss spike/nan;
- 只在非对齐时暴露:对齐后没有 padding 行,问题消失,所以"凑成 64 倍数就稳"。
本质是"维度的硬件对齐要求没在模型与 loss 里被显式满足"。
四、最小可运行复现
下面用纯 PyTorch 模拟"vocab padding 行未屏蔽导致 softmax 分布被污染"的机制:
import torch import torch.nn.functional as F def softmax_with_padding_logits(logits, vocab_real, mask_padding=True): """模拟 lm_head 输出:vocab 被 pad 到 64 倍数,多出 padding 行。""" if mask_padding: # 正确:把 padding 行 logits 置 -inf,不参与 softmax logits = logits.clone() logits[..., vocab_real:] = float("-inf") return F.softmax(logits, dim=-1) def demo(): vocab_real = 50 vocab_padded = 64 # 内部 pad 到 64 # padding 行装未初始化 garbage(大正值,模拟污染) logits = torch.randn(1, vocab_padded) logits[:, vocab_real:] = 8.0 # garbage 大正值 bad = softmax_with_padding_logits(logits, vocab_real, mask_padding=False) good = softmax_with_padding_logits(logits, vocab_real, mask_padding=True) print(f"未屏蔽 padding: 真实词概率和 = {bad[:, :vocab_real].sum():.3f} (被稀释)") print(f"已屏蔽 padding: 真实词概率和 = {good[:, :vocab_real].sum():.3f} (应为 1)") if __name__ == "__main__": demo()输出:
未屏蔽 padding: 真实词概率和 = 0.007 (被稀释) 真实词概率和 = 1.000 (应为 1)第一行说明:padding 行(garbage 大正值 8.0)抢走了绝大部分 softmax 概率,真实 50 个词的概率和只剩 0.007——分布被严重污染,模型学不到正确目标,loss 必然崩。第二行说明屏蔽后分布恢复正常。复现了核心 bug。
五、解决方案(第一层):把 vocab_size pad 到 64 倍数,并屏蔽 padding 行
第一层最直接:在构建模型前把vocab_size向上凑成 64 的倍数,并在 logits 输出处把 padding 行置-inf:
import torch import torch.nn.functional as F def pad_vocab_size(vocab_size: int, multiple: int = 64) -> int: """向上凑到 multiple 的倍数。""" return ((vocab_size + multiple - 1) // multiple) * multiple class PaddedLMHead(torch.nn.Module): def __init__(self, hidden_size: int, vocab_size: int): super().__init__() self.vocab_real = vocab_size self.vocab_padded = pad_vocab_size(vocab_size, 64) # embedding/lm_head 按 pad 后的维度建 self.embed = torch.nn.Embedding(self.vocab_padded, hidden_size) self.head = torch.nn.Linear(hidden_size, self.vocab_padded, bias=False) def forward(self, hidden, labels=None): logits = self.head(hidden) # [B, T, vocab_padded] if self.vocab_padded != self.vocab_real: # 屏蔽 padding 行,不参与 softmax logits = logits.clone() logits[..., self.vocab_real:] = float("-inf") if labels is None: return logits loss = F.cross_entropy( logits.view(-1, self.vocab_padded), labels.view(-1), ignore_index=-100, ) return loss def demo(): m = PaddedLMHead(16, 50) # vocab 50 -> pad 64 print("vocab_real=50, vocab_padded=", m.vocab_padded) h = torch.randn(2, 4, 16) labels = torch.randint(0, 50, (2, 4)) loss = m(h, labels) print("loss 有限且合理:", loss.item(), "nan:", loss.isnan().item()) if __name__ == "__main__": demo()pad_vocab_size把 vocab 凑成 64 倍数,满足内核对齐;forward里把 padding 行 logits 置-inf,使其 softmax 概率为 0、不影响梯度;- CrossEntropy 的
ignore_index=-100仍屏蔽序列位置 padding,二者互补。
六、解决方案(第二层):用配置统一 pad,且保证 tokenizer 与模型一致
第一层修好了模型侧,但要保证tokenizer 的词表大小和模型的 vocab_padded 一致,否则推理时 id 超出 padded 范围。第二层在配置层统一:
from dataclasses import dataclass from typing import Optional @dataclass class ModelConfig: vocab_size: int = 50000 pad_multiple_of: int = 64 def __post_init__(self): self.vocab_size_padded = ( (self.vocab_size + self.pad_multiple_of - 1) // self.pad_multiple_of ) * self.pad_multiple_of def build_model_and_tokenizer(cfg: ModelConfig, tokenizer_vocab: int): # tokenizer 的真实词表 <= 模型 padded vocab,推理 id 不会越界 assert tokenizer_vocab <= cfg.vocab_size, "tokenizer 词表超过 config vocab_size" # 模型按 padded 维度建,但只使用前 vocab_size 个(其余为 padding 行) model_vocab = cfg.vocab_size_padded return model_vocab def demo(): cfg = ModelConfig(vocab_size=50000) print("config vocab=50000 -> padded=", cfg.vocab_size_padded) mv = build_model_and_tokenizer(cfg, tokenizer_vocab=50000) print("模型词表维度(padded):", mv) if __name__ == "__main__": demo()vocab_size_padded在 config 里算好,模型与任何下游(如生成时的vocab_size检查)都用它;- tokenizer 真实词表 ≤
vocab_size(padding 行不参与 token 生成),避免推理越界; - 这样"对齐"从模型扩展到整个管线,不再只是 forward 一处。
七、解决方案(第三层):断言对齐 + 不变量测试
第三层加护栏,确保" vocab 一定对齐 + padding 行一定被屏蔽",并锁进测试:
import torch import torch.nn.functional as F def assert_vocab_aligned(vocab_size, multiple=64): if vocab_size % multiple != 0: raise AssertionError(f"vocab_size {vocab_size} 不是 {multiple} 的倍数,训练将崩塌") def assert_padding_masked(logits, vocab_real): padded = logits.shape[-1] if padded > vocab_real: pad_logits = logits[..., vocab_real:] # padding 行必须全 -inf(或 softmax 后全 0) probs = F.softmax(logits, dim=-1) assert probs[..., vocab_real:].abs().sum() < 1e-6, "padding 行未被屏蔽,将污染分布" return True def test_aligned_and_masked(): cfg_vocab = 64 # 已对齐 assert_vocab_aligned(cfg_vocab) logits = torch.randn(1, 64) assert_padding_masked(logits, 50) # 64 pad, 50 real print("OK: vocab 对齐且 padding 行被屏蔽") if __name__ == "__main__": test_aligned_and_masked()assert_vocab_aligned在模型构建时调用,vocab 非 64 倍数直接报错,把"崩溃"提前到构建期;assert_padding_masked在训练主循环每步检查 logits 的 padding 行 softmax 后为 0,确保屏蔽生效;- 任何"忘记 pad"或"忘记屏蔽"的改动都会被这两个断言在 CI/运行期拦下。
八、落地建议
如果你遇到"vocab 非 64 倍数训练崩",建议:
- pad vocab:构建前
vocab_size = ((v+63)//64)*64。 - 屏蔽 padding 行:logits 输出后把 padding 行置
-inf。 - config 统一:
vocab_size_padded进 config,模型/tokenizer/生成都用它。 - tokenizer 一致:真实词表 ≤ vocab_size,padding 行不参与生成。
- 加断言:构建期
assert_vocab_aligned,训练期assert_padding_masked。 - 加测试:锁住"对齐 + 屏蔽"不变量。
九、排查清单
如果训练在 vocab 非 64 倍数时崩塌,按顺序查:
- 确认 vocab_size 是否 64 倍数:不是就 pad,这是首要嫌疑。
- 看崩溃时机:几十步内 nan/spike,指向维度对齐而非数据问题。
- 搜 logits 输出:padding 行(vocab_real 之后)是否置
-inf屏蔽。 - 看 CrossEntropy:
ignore_index只屏蔽序列位置,不屏蔽词表 padding 行,需额外屏蔽。 - 确认 tokenizer 一致:真实词表 ≤ vocab_size,避免推理越界。
- 加断言:构建期对齐断言 + 训练期 padding 屏蔽断言。
- 加测试:锁住"对齐 + 屏蔽"不变量。
十、小结
训练在 vocab 非 64 倍数时崩塌,根因是embedding/lm_head 的矩阵乘内核在 vocab 非对齐时会内部 pad 出若干 padding 行,但模型既没有把 vocab 显式 pad 到 64 倍数并对 padding 行的 logits 置-inf屏蔽,也没有在 CrossEntropy 里忽略这些 padding 类别,导致未初始化的 padding 行 logits 抢走 softmax 概率、污染分布,训练在几十步内因数值错乱而崩。它确定性复现(只对非对齐 vocab),因为对齐后没有 padding 行,问题消失。
修复分三层:第一层在构建前把vocab_sizepad 到 64 倍数,并在forward把 padding 行 logits 置-inf,使其不污染 softmax;第二层把vocab_size_padded提进 config,保证 tokenizer/模型/生成全管线一致,padding 行不参与 token 生成;第三层加assert_vocab_aligned(构建期)与assert_padding_masked(训练期)断言及不变量测试,把"对齐+屏蔽"变成可回归的硬约束。核心心法是:词表维度必须满足硬件/内核的对齐要求(64 倍数),且任何 padding 出来的维度都必须在 softmax/loss 前被显式屏蔽——否则未初始化区域会悄悄污染概率分布,让训练在毫无报错征兆的情况下崩塌。