已合并
fix: KimiK3 use fused npu_rms_norm in KimiRMSNorm, fix chunk_kda. #2899
suis-moi创建于 19 天前
fix: KimiK3 use fused npu_rms_norm in KimiRMSNorm, fix chunk_kda. #2899
已合并
suis-moi创建于 19 天前
suis-moi
suis-moi成员
19 天前

What this PR does / why we need it?

1、重构MOE模块
2、更新开源最新版代码
3、修复chunk kda.py
4、替换kimi RMSNORM为融合算子

Does this PR introduce any user-facing change?

Please describe whether the PR will result in any user-facing usage changes. If there is related documentation, please specify its path.

How was this patch tested?

本地已经验证通过

likedislike
Pull Request已成功合入, 合并人@ascend-robot
(感谢 suis-moi 的贡献)
suis-moisuis-moi成员
19 天前 创建了 pull request,commit b8fdb1cb
atomgit-bot
atomgit-bot
19 天前 评论:

变更摘要

该 PR 主要围绕两个目标:在 NPU 设备上为 KimiRMSNorm 使用融合算子 torch_npu.npu_rms_norm 以提升性能,以及对 chunk_kda 算子进行大规模重构——移除已弃用的参数(use_beta_sigmoid_in_kernelallow_neg_eigvalstate_v_firstcp_contextchunk_size),同时新增 skip_recompute 机制,在梯度检查点重放阶段通过主机端卸载/恢复中间结果来跳过前向重计算,节省显存与计算开销。此外,将 NPU 相关的导入和逻辑收敛到 IS_NPU_AVAILABLE 条件分支下,使代码在 NPU 与 GPU 环境下均能正确加载对应依赖。

主要改动

  • KimiRMSNorm 改用 NPU 融合算子: 在 NPU 上直接调用 torch_npu.npu_rms_norm 替代原先手动 float() + rsqrt 的实现,消除了 dtype 转换开销并利用硬件加速。
  • chunk_kda 移除多个已弃用/冗余参数: 删除了 use_beta_sigmoid_in_kernelallow_neg_eigvalstate_v_firstcp_contextchunk_size 等参数及相关逻辑(包括 fused_beta_sigmoid/fused_beta_sigmoid_bwd),简化了对外接口与内部实现。
  • 新增 skip_recompute 内存优化机制: 在前向阶段通过 OffloadManageroAqkAkk 等中间张量卸载到主机内存;在反向阶段重放时从主机恢复,避免重新执行 chunk_kda_fwd,实现比特一致梯度且节省计算。
  • 导入路径重构: chunk.py 中将对 fla 库的依赖替换为本地模块(如 .l2norm_kda.chunk_bwd.chunk_fwd.utils.fla_utils),modeling_kimi_linear.py 中将 NPU 相关导入收敛到 IS_NPU_AVAILABLE 条件分支内。
  • 参数重命名: 将 transpose_state_layout 从 deprecated 的别名正式提升为独立参数,替代原有的 state_v_first
likedislike
atomgit-bot
atomgit-bot
19 天前 评论:

代码审查

审查总结

本次 diff 涉及 5 个文件, 已全部审查完毕:

文件 审查结论
Third-Party Open Source Software Notice.txt 无问题 — 仅新增 Kimi K3 许可证文本
examples/kimi_k3/README.md 无问题 — 目录重构 + Triton-Ascend 安装说明
mindspeed_mm/fsdp/models/kimi_k3/LICENSE 无问题 — 新增许可证文件
mindspeed_mm/fsdp/models/kimi_k3/modeling_kimi_linear.py 1 个 P3 问题(残留无效参数)
mindspeed_mm/fsdp/ops/kda/triton_ascend/chunk.py 1 个 P2 + 1 个 P3 问题

按优先级统计:

  • P2:1 个 — skip_recompute 路径中 g_cumsum=None 传入 chunk_kda_bwd, 若后端不支持 None 会崩溃(置信度 0.5, 因无法验证 chunk_bwd.py
  • P3:2 个 — 调用方残留无效参数 use_beta_sigmoid_in_kernel=Falsechunk_kda**kwargs 静默吞没废弃参数

整体风险评估:中低。 核心运行时路径(非 skip_recompute 的正常训练、NPU RMSNorm 融合算子、KimiDeltaAttention 调用)均正确无误。skip_recompute 特性作为新增功能, 其设计意图清晰(不卸载 g_cumsum 表明后端应能从 g_org 重建), 但缺少该后端代码的可见性, 建议在集成环境验证该路径。P3 的清理项不影响正确性, 可在后续迭代中处理。

⚠️ 已识别出整体风险,但无法提取行内评论,请参考整体评估。

likedislike
ascend-robotascend-robot成员
19 天前 添加了label:stat/needs-squash
ascend-robotascend-robot成员
19 天前 添加了label:ascend-cla/yes
此处折叠了131条消息 查看更多
suis-moisuis-moi成员
16 天前 修改了pull request 的描述
htwang成员
16 天前 评论:

同意合入

likedislike
htwang成员
16 天前 评论:

/approve

likedislike
ascend-robotascend-robot成员
16 天前 添加了label:approved
ascend-robotascend-robot成员
16 天前 合入了pull request