【Bug已解决】Training collapses when token embeddings are not padded to a multiple of 64 解决方案

【Bug已解决】Training collapses when token embeddings are not padded to a multiple of 64 解决方案 【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: loss2.1 (正常) step 35: lossnan # 或 step 40: loss87.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 了 vocabCrossEntropy 的ignore_index只屏蔽了序列位置的 paddingtoken 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_paddingTrue): 模拟 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_paddingFalse) good softmax_with_padding_logits(logits, vocab_real, mask_paddingTrue) 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 行置-infimport 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, biasFalse) def forward(self, hidden, labelsNone): 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_real50, 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_size50000) print(config vocab50000 - padded, cfg.vocab_size_padded) mv build_model_and_tokenizer(cfg, tokenizer_vocab50000) print(模型词表维度padded, mv) if __name__ __main__: demo()vocab_size_padded在 config 里算好模型与任何下游如生成时的vocab_size检查都用它tokenizer 真实词表 ≤vocab_sizepadding 行不参与 token 生成避免推理越界这样对齐从模型扩展到整个管线不再只是 forward 一处。七、解决方案第三层断言对齐 不变量测试第三层加护栏确保 vocab 一定对齐 padding 行一定被屏蔽并锁进测试import torch import torch.nn.functional as F def assert_vocab_aligned(vocab_size, multiple64): if vocab_size % multiple ! 0: raise AssertionError(fvocab_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 ((v63)//64)*64。屏蔽 padding 行logits 输出后把 padding 行置-inf。config 统一vocab_size_padded进 config模型/tokenizer/生成都用它。tokenizer 一致真实词表 ≤ vocab_sizepadding 行不参与生成。加断言构建期assert_vocab_aligned训练期assert_padding_masked。加测试锁住对齐 屏蔽不变量。九、排查清单如果训练在 vocab 非 64 倍数时崩塌按顺序查确认 vocab_size 是否 64 倍数不是就 pad这是首要嫌疑。看崩溃时机几十步内 nan/spike指向维度对齐而非数据问题。搜 logits 输出padding 行vocab_real 之后是否置-inf屏蔽。看 CrossEntropyignore_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 前被显式屏蔽——否则未初始化区域会悄悄污染概率分布让训练在毫无报错征兆的情况下崩塌。