已开启
[Bug]: fully_shard replicate_params 的 BF16 临时参数覆盖 FP32 主参数 #303
MengXY107创建于  7月24日
MengXY107
MengXY107成员
7月24日 创建

Checklist

🐛 Describe the bug

fully_shard 管理的 replicate_params 以 FP32 初始化,并通过 MixedPrecisionPolicy(param_dtype=torch.bfloat16) 使用 BF16 参与前向计算时,Torch 与 MindSpore 两端的 HSDPParam.to_sharded() 都会把临时 unsharded 参数复制回 sharded 主参数。

由于临时参数已被转换成 BF16,这次复制会将 BF16 数值写回 FP32 sharded_param。虽然目标张量的 dtype 仍显示为 FP32,但参数有效精度已永久舍入为 BF16;后续优化器更新不再基于原始 FP32 主参数。

最小复现流程:

model = ExistingModel().to("npu")  # 参数初始化为 FP32
replicate_params = set(model.target.parameters())
fully_shard(
    model,
    mesh=mesh,
    reshard_after_forward=True,
    mp_policy=MixedPrecisionPolicy(param_dtype=torch.bfloat16),
    replicate_params=replicate_params,
)
master_before = [param.detach().clone() for param in model.target.parameters()]
model(inputs)
# 当前行为:forward/reshard 后 FP32 主参数被 BF16 临时参数覆盖

该问题是数值精度劣化,不会产生 Python traceback。

Expected behavior

to_sharded() 只负责恢复 module 上的 sharded 参数对象并释放 unsharded 临时存储,不应把低精度临时参数反向覆盖到 FP32 主参数。前向/reshard 后,replicate_params 的 FP32 主参数 dtype 与数值都应保持不变,优化器继续基于 FP32 主参数更新。

Additional context

Torch 与 MindSpore 实现存在相同复制逻辑,需要保持两端语义一致。回归验证复用现有 Torch replicate_params 精度场景:standalone 网络整体转换为 BF16,fully_shard 网络仍以 FP32 初始化并设置 param_dtype=BF16;不新增测试用例或 module 类。

Environment info

  • Repository: mindspore/hyper-parallel
  • Base: upstream/master at bb62140e
  • Backends: PyTorch and MindSpore
  • Target device: Ascend NPU
likedislike
MengXY107MengXY107成员
7月24日 添加了label:bug
MengXY107MengXY107成员
7月24日 关联了pull request:fix: 避免 replicate_params 的 FP32 主参数精度劣化