已合并
[bugfix][pynative] 将 MoE/MTP/index loss 追踪重构为模型实例方法并修复 recompute 重复累加问题 #8465
[bugfix][pynative] 将 MoE/MTP/index loss 追踪重构为模型实例方法并修复 recompute 重复累加问题 #8465
已合并
niujunhao创建于 6月25日
niujunhao成员
6月25日

https://gitcode.com/mindspore/org-issues/issues/42712
https://gitcode.com/mindspore/org-issues/issues/42716

修改描述 (Description)

背景

在 TP×CP 并行下,MoE 路由的 token 视图被切分到多个 rank 上:

  • TP 维度:router 输入在 TP rank 间被复制(路由前未做 token 切分),但其后的 dispatch 阶段会切分
  • CP 维度:context parallel 把 sequence 切到不同 rank

旧实现中,TopKRouter 直接用本地 routing_map.sum(dim=0) 计算 tokens_per_expert,没有在 TP×CP 组内做 all-reduce 求和。这导致:

  1. aux-loss 计算时每个 rank 只看到自己那份 token 的直方图,与单卡基线不一致
  2. total_num_tokens 用的是 seq_length(local),没有还原成全局 token 数
  3. 在 recompute 触发时,MoE 层的 tokens_per_expert 会被累加两次,aux-loss 随步数线性漂移
  4. expert_bias 更新用的是本地直方图,bias delta 与单卡不一致
  5. MTP 损失在 recompute 时也会被 tracker 累加两次
  6. LossCallback 直接调用 tracker 函数,耦合度高,无法适配非 GPTModel 模型

修复

核心修改 — 在 TP×CP 组内做正确的 all-reduce:

  • 新增 moe_utils.get_tokens_per_expert_and_token_count(routing_map, reduce_group, topk, with_padding_mask):
    • 用 routing_map.sum(dim=0) 算 local per-expert 直方图
    • 沿 reduce_group(仅 TP×CP,不含 DP)做 in-place SUM all-reduce,还原 global 直方图
    • 推导 total_num_tokens = local_num_tokens * group.world_size()
    • DP 故意排除:DP 各 rank 本就该分别贡献自己的 aux-loss,由调用方用 /bsz 做平均;如果对 DP 做 all-reduce 会重复计数
  • 新增 moe_utils.get_world_size_from_group(group),兼容 hyper_parallel / 裸 group / get_world_size 三种 group size 获取方式,失败时降级为 1

TopKRouter 改造:

  • 新增实例属性 self.aux_loss_group = None、self.aux_loss_group_size = 1,由外部在并行模式下注入
  • seq_aux_loss 和 global_aux_loss 两条路径都改为调用 get_tokens_per_expert_and_token_count(...),去掉手写的 tokens_per_expert = routing_map.sum(dim=0) 和手算的
    total_num_tokens
  • 注释里说明 topk * bsz 是为了和 Megatron helper 签名对齐(当前未启用 padding 分支)

Recompute 双计防护:

  • MoELayer.forward:在 tokens_per_expert.add_(...) 之前加 if not is_in_recompute(): 守卫,避免重计算阶段把 token 数累加两遍
  • multi_token_prediction.save_to_mtp_losses_tracker:同样的 is_in_recompute() 守卫,避免 MTP 损失在 tracker 里被重复累加

Expert bias 更新修复:

  • GPTModel._update_expert_bias(metric_group, metric_group_size):
    • 在 metric_group(dp×cp 域)上对每层 tokens_per_expert 做 SUM all-reduce,恢复成 global-batch 直方图
    • DP 必须包含:每张 DP 卡看到的是 global batch 的不同 shard,本地直方图只是部分计数,不 reduce 就会和单卡基线发散
    • TP 故意排除:MoE plan 把 router 输入在 TP 间复制,TP 看到同样的 token,对 TP 归约会把计数乘以 tp_size
    • 用 _no_grad() 上下文执行 bias 更新和 tokens_per_expert 归零
  • 引入 _get_expert_bias_modules() / _get_global_aux_loss_modules(),用实例属性缓存模块列表,VPP 下每个 chunk 一份独立 cache,互不干扰

模型接口泛化(LossCallback 解耦):

  • GPTModel 暴露 5 个方法取代 callback 直接调用 tracker:
    • _update_expert_bias(metric_group, metric_group_size)
    • get_load_balancing_loss(...) — 调用 track_moe_metrics(PP-group 组合 + dp/cp 域归约),每 stage 每 step 都必须调用以避免 deadlock
    • reset_model_temporary_tensors() — global_aux_loss 模式下重置 tracker
    • get_mtp_loss(metric_group, metric_group_size) — 调用 track_mtp_metrics
    • get_index_loss() — 调用 track_indexer_metrics
  • PyNativeDeepseekV3ForCausalLM 透传这 5 个方法到内部 self.model
  • LossCallback 改为通过 hasattr 探测方法是否存在并调用,缺失时只 warn 一次(用 logger_record 字典去重),错误信息附带具体模型类名

行为变化

  • TopKRouter 现在要求外部在并行模式下注入 aux_loss_group;单卡/未启用并行的场景保持 None,行为与旧实现一致
  • aux-loss 不再随训练步数漂移(TP×CP + DP 都正确还原)
  • expert_bias 增量与单卡基线一致
  • 触发 recompute 时 MoE token 计数和 MTP loss 不再被双重累加
  • LossCallback 不再硬编码 GPTModel,可被其它 TrainModelMixin 模型复用

修改类型 (Type of Change)

测试结果 (Test Results)

image.png
image.png

检查清单 (Checklist)

likedislike
Pull Request已成功合入, 合并人@MindSpore-Bot
(感谢 niujunhao 的贡献)
Nniujunhao成员
6月25日 创建了 pull request,commit 8d1aebf9
MindSpore-BotMindSpore-Bot成员
6月25日 添加了label:mindspore-cla/yes
司小南(机器人)
司小南(机器人)成员
6月25日 评论:

@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触发检查。

likedislike
司小南(机器人)司小南(机器人)成员
6月25日 添加了label:pr-check-fail
atomgit-bot
atomgit-bot
6月25日 评论:

变更摘要

此 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 独立缩放,不受主线损失缩放影响。
likedislike
不准确?
此处折叠了168条消息 查看更多
MindSpore-Bot
MindSpore-Bot成员
6月30日 评论:

Review Code Feedback

  • The label lgtm-hss-shuai, approved was added to this pull request. It means that hss-shuai reviewed the code changes. 👋
Tips
  • If this pull request is not merged while all conditions are met, comment /check-pr to try again. 😄
likedislike
husichao成员
6月30日 评论:

/lgtm

likedislike
MindSpore-BotMindSpore-Bot成员
6月30日 添加了label:lgtm-husichao
MindSpore-Bot
MindSpore-Bot成员
6月30日 评论:

Review Code Feedback

  • The label lgtm-husichao was added to this pull request. It means that husichao reviewed the code changes. 👋
Tips
  • If this pull request is not merged while all conditions are met, comment /check-pr to try again. 😄
likedislike
MindSpore-BotMindSpore-Bot成员
6月30日 合入了pull request,合并节点 SHA:5c328d8ded0083e5609aa49d0bad7a9be96576ae