Wrong router gradient with `experts_implementation=”batched_mm”` under expert parallelism

在专家并行(EP)下使用 experts_implementation="batched_mm" 时,MoE 路由门控权重( mlp.gate.weight )会收到错误梯度,loss 与参考实现完全一致,但梯度偏大约 50%;优先确认是否启用了 batched_mm ,并切换为 grouped_m

快速结论:在专家并行(EP)下使用 experts_implementation="batched_mm" 时,MoE 路由门控权重(mlp.gate.weight)会收到错误梯度,loss 与参考实现完全一致,但梯度偏大约 50%;优先确认是否启用了 batched_mm,并切换为 grouped_mmeager 做对照。

适用环境:Issue 中确认的环境为:transformers main(commit c694707483)、PyTorch 2.9、2x H100 + NCCL,另外在 CPU + gloo 上也可复现。

最快修复方案:暂无确认的一步修复方案。Issue 中明确指出按源码分析 batched_mm 缺少 sentinel 掩码,并称“PR coming”,但讨论链未给出已合并的修复版本或补丁。可优先尝试把 experts_implementation 改为 grouped_mm(或 eager)规避该问题。

注意事项:该问题不影响前向 loss,只影响反向梯度,因此普通训练日志不会暴露异常。切换 experts_implementation 只是规避手段,是否有性能差异、是否影响其他并行配置,Issue 未给出验证结论。

问题场景

用户在 expert parallelism(EP)模式下,用 OlmoeForCausalLM 加载模型并设置 set_experts_implementation("batched_mm") 后做训练/反向传播,发现 router 梯度与单卡或非 EP 参考实现不一致。Issue 中提到用 torchrun --nproc_per_node 2 运行复现脚本,分别对比启用 EP 的模型和作为参考的非 EP 模型,比较 model.layers.*.mlp.gate.weight 的梯度范数。

报错原文

Wrong router gradient with `experts_implementation="batched_mm"` under expert parallelism

router grad norm: reference 2.240401e-04  EP 2.846920e-04

原因分析

最可能的原因是 batched_mm_experts_forward 在处理 EP sentinel slot 时只 clamp 了 expert_ids,并依赖对应 routing weight 为零来保证前向不贡献结果,但没有在前向中把 sentinel slot 的 proj_out 置零。

前向看起来是正确的:weight 为零,所以 sentinel slot 对输出没有贡献。但反向时,对 sample_weights 求导得到的是 proj_out,而 sentinel slot 的 proj_out 并不为零(实际是 expert 0 在该 token 上的计算结果),于是本 rank 从未路由到的 slot 也会把梯度灌进 top_k_weights,再传到 router。Issue 作者对比指出,同文件中的 grouped_mm_experts_forward 已在两次 matmul 之前用 sentinel_maskproj_out 置零,而 batched_mm 缺了这一步。

另外,作者指出当前 EP 测试无法捕获该问题,因为 _test_ep_backward_impl 只比较 loss,而 loss 不受影响;#48518 改为比较逐参数梯度后才暴露出来。

环境排查

  • 确认 transformers 版本:Issue 使用 main 分支 commit c694707483。
  • 确认 PyTorch 版本:Issue 环境为 torch 2.9。
  • 确认 GPU 与通信后端:2x H100 + NCCL;CPU + gloo 也可复现。
  • 确认是否启用 expert parallelism:DistributedConfig(tp_size=world, enable_expert_parallel=True)
  • 确认 experts implementation:是否设置为 batched_mm;可切换为 grouped_mmeager 做对照。
  • 确认模型层数、专家数、num_experts_per_tok 等配置与复现脚本一致。

解决步骤

  1. 先确认当前代码是否使用 experts_implementation="batched_mm";如果是,优先尝试改为 grouped_mmeager,观察 router 梯度是否与参考实现一致。
  2. 准备一个非 EP 参考模型,并让两个模型都执行 forward + backward,比较 model.layers.0.mlp.gate.weight 的梯度范数。
  3. 检查其他参数梯度是否一致:Issue 中除 mlp.gate.weight 外,其余参数误差在 1e-5 以内。
  4. 如果必须使用 batched_mm,请关注该 Issue 后续 PR 是否合入;在修复合入前不要仅凭 loss 判断训练正确性。
  5. 若正在做 EP 相关测试或回归,建议增加逐参数梯度对比,而不是只比较 loss。

验证方法

运行复现脚本,比较 EP 模型与参考模型的 router 梯度范数。Issue 中参考值为 reference 2.240401e-04、EP 为 2.846920e-04,同时 loss 完全一致,其他参数一致到 1e-5。若切换为 grouped_mmeager 后 router 梯度与参考一致,而 batched_mm 仍偏大,即可确认问题与 batched_mm 相关。

参考来源

huggingface/transformers #48687

GamsGo AI

AI 工具推荐

想把多个 AI 模型放在一个入口?

GamsGo AI 集成 ChatGPT、DeepSeek、Gemini、Claude、Midjourney、Veo 等常用模型,适合写作、绘图、视频和日常 AI 工作流。

了解 GamsGo AI

推广链接:通过此链接购买,我可能获得佣金,不影响你的价格。

这个方案解决了吗?

celebrityanime
celebrityanime
文章: 23059

发表回复

您的邮箱地址不会被公开。 必填项已用 * 标注