已开启
[Bug]: PyNative Muon 首步 master weight 回拷临时显存延迟释放 #2530
Java不加糖创建于  8月19日
Java不加糖成员
8月19日 创建

Checklist

🐞 问题详细描述

PyNative Muon 在 BF16/FP16 模型参数上使用 FP32 master weight。优化器更新结束后,master-to-model 同步当前执行:

inplace_copy(model_param, op_cast(fp32_param, model_param.dtype))

显式 op_cast 会物化一个完整模型参数大小的低精度临时 Tensor。该 Tensor 的异步拷贝生命周期可能延伸到下一次前向边界,导致首步优化器更新后的显存未及时回落。该现象最初在大规模 TeleChat/DeepSeekV4 BF16 Muon 训练的 memory tracker 中发现。

最小复现使用真实 Ascend NPU,按生产 Trainer 的方式将 optimizer(grads) 作为裸表达式调用,不保存 optimizer 返回值。512×512 BF16 参数下:

  • 显式 Cast + InplaceCopy:额外产生一个 524,800 B 非持久显存块;
  • FP32 master 直接跨 dtype InplaceCopy:对应显存块消失;
  • BF16、FP16 数值结果均与原路径 bitwise equal。

Muon 的 construct 返回值在 PyNative Trainer 中未使用,但该返回值不是本问题根因;修复不改变其返回契约。

期望 master weight 回拷不物化完整参数大小的显式 Cast Tensor,并保持现有数值结果。

详细的环境信息描述

  • Hardware:Ascend 910B2
  • MindSpore:2.10
  • CANN:9.1.0-beta.3
  • Mode:PyNative
  • Optimizer:Muon
  • Parameter dtype:BF16/FP16
  • Master weight dtype:FP32
  • Communication strategy:allgather

其他辅助信息

建议使用 memory tracker 比较首次 optimizer update 到第二次 forward 边界的非持久显存块,并覆盖 BF16、FP16 跨 dtype 回拷数值一致性。

版本信息

  • r2.1.0-beta1
  • master
likedislike
JJava不加糖成员
8月19日 关联了pull request:fix(pynative): avoid Muon master-copy cast temporary