Pull Request已成功合入, 合并人@MindSpore-Bot
(感谢 niujunhao 的贡献)@alpha-junh, 当前/check-pr未通过,原因如下:
以下Pull Request描述检查项未通过:
存在不符合模板的选项: Bug 修复 (Bug fix)
存在不符合模板的选项: 新特性 (New feature)
存在不符合模板的选项: 文档更新 (Documentation update)
存在不符合模板的选项: 代码重构 (Refactoring)
存在不符合模板的选项: 性能优化 (Performance improvement)
存在不符合模板的选项: 其他 (Other)
存在不符合模板的选项: 我已自验过该功能/修复 (Self-checked)
存在不符合模板的选项: 我已通过了本地的单元测试 (Passed local UTs)
存在不符合模板的选项: (如适用) 我已更新了对应文档 (Updated documentation)
存在不符合模板的选项: (如适用) 我已添加了新的测试用例 (Added new tests)
部分检查项缺失 请重新使用模板
模板中'Test Plan and Test Result' 信息为空,请补充对应信息。
以下issue检查项未通过:
Pull Request未关联issue
请修改好上述检查错误后,重新使用/check-pr触发检查。


变更摘要
此 PR 主要将 MoE 辅助损失、MTP 损失及 indexer 损失的追踪逻辑从 LossCallback 回调中提取为 GPTModel 的实例方法,并修复了 recompute 场景下 tokens_per_expert 累加和 MTP loss 被重复计数导致精度偏差的问题。重构后,LossCallback.on_step_end 通过调用模型实例方法获取各类损失,代码逻辑更清晰;同时在 moe_layer.py 和 multi_token_prediction.py 中增加了 is_in_recompute() 守卫,避免重计算前向路径的重复累加。
主要改动
- 损失追踪重构为
GPTModel实例方法:在GPTModel中新增_update_expert_bias、get_load_balancing_loss、get_mtp_loss、get_index_loss、reset_model_temporary_tensors五个方法,封装原本散落在LossCallback中的 MoE/MTP/index 损失汇聚与归约逻辑,并通过track_moe_metrics、track_mtp_metrics、track_indexer_metrics等工具函数完成跨 rank 通信。 LossCallback大幅简化:on_step_end从原先内联的损失处理逻辑(约 134 行)改为仅调用模型实例方法,减少了回调与模型内部实现的耦合。- 修复 recompute 重复累加:在
moe_layer.py的tokens_per_expert.add_前和multi_token_prediction.py的save_to_mtp_losses_tracker中增加is_in_recompute()检查,当处于重计算前向时跳过累加,避免梯度累积步内重复计入导致数值偏差。 PyNativeDeepseekV3ForCausalLM增加委托方法:为 DeepSeek-V3 模型添加_update_expert_bias、get_load_balancing_loss、reset_model_temporary_tensors、get_mtp_loss、get_index_loss五个方法,将调用转发到内部self.model(即GPTModel实例),使回调可通过hasattr统一适配。- PP 损失缩放注释更新:在
parallelize.py的 PP 损失缩放逻辑中新增注释说明 MoE/MTP 损失由各自的 auto-scaler 独立缩放,不受主线损失缩放影响。


关联 Issue (Related Issue)
https://gitcode.com/mindspore/org-issues/issues/42712
https://gitcode.com/mindspore/org-issues/issues/42716
修改描述 (Description)
背景
在 TP×CP 并行下,MoE 路由的 token 视图被切分到多个 rank 上:
旧实现中,TopKRouter 直接用本地 routing_map.sum(dim=0) 计算 tokens_per_expert,没有在 TP×CP 组内做 all-reduce 求和。这导致:
修复
核心修改 — 在 TP×CP 组内做正确的 all-reduce:
TopKRouter 改造:
total_num_tokens
Recompute 双计防护:
Expert bias 更新修复:
模型接口泛化(LossCallback 解耦):
行为变化
修改类型 (Type of Change)
测试结果 (Test Results)
检查清单 (Checklist)