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 中发现。
op_cast
最小复现使用真实 Ascend NPU,按生产 Trainer 的方式将 optimizer(grads) 作为裸表达式调用,不保存 optimizer 返回值。512×512 BF16 参数下:
optimizer(grads)
Muon 的 construct 返回值在 PyNative Trainer 中未使用,但该返回值不是本问题根因;修复不改变其返回契约。
construct
期望 master weight 回拷不物化完整参数大小的显式 Cast Tensor,并保持现有数值结果。
建议使用 memory tracker 比较首次 optimizer update 到第二次 forward 边界的非持久显存块,并覆盖 BF16、FP16 跨 dtype 回拷数值一致性。
r2.1.0-beta1
master
Checklist
🐞 问题详细描述
PyNative Muon 在 BF16/FP16 模型参数上使用 FP32 master weight。优化器更新结束后,master-to-model 同步当前执行:
显式
op_cast会物化一个完整模型参数大小的低精度临时 Tensor。该 Tensor 的异步拷贝生命周期可能延伸到下一次前向边界,导致首步优化器更新后的显存未及时回落。该现象最初在大规模 TeleChat/DeepSeekV4 BF16 Muon 训练的 memory tracker 中发现。最小复现使用真实 Ascend NPU,按生产 Trainer 的方式将
optimizer(grads)作为裸表达式调用,不保存 optimizer 返回值。512×512 BF16 参数下:Muon 的
construct返回值在 PyNative Trainer 中未使用,但该返回值不是本问题根因;修复不改变其返回契约。期望 master weight 回拷不物化完整参数大小的显式 Cast Tensor,并保持现有数值结果。
详细的环境信息描述
其他辅助信息
建议使用 memory tracker 比较首次 optimizer update 到第二次 forward 边界的非持久显存块,并覆盖 BF16、FP16 跨 dtype 回拷数值一致性。
版本信息
r2.1.0-beta1master