在 FSDP/HSDP(单机 8 卡,DP=8)下,clip_grad_norm_ 计算全局梯度范数时,把非切分参数(replicate 参数,或经 FSDP ignored_params 由用户自管理的参数)的本地 norm² 也放进了 shard 进程组的 all_reduce(SUM)。这些梯度本身在各 rank 上已是全局一致的值,其 norm² 应只在本地计入一次;多做一次 shard 组 SUM 会把它重复计 shard_world_size 次(8 卡即 ×8),使全局 grad_norm 偏大。
clip_grad_norm_
ignored_params
all_reduce(SUM)
shard_world_size
两种管理不切参数的方式都复现同一症状(并行策略未改、首步 loss 未变,仅 grad_norm 改变):
replicate_params
根因:replicate / ()-签名 的梯度应"本地计入、不通信",旧实现却在 shard 组上对其 norm² 求和。注:首步 loss 不变是因为该步 norm 未超过 max_norm、未触发实际裁剪,但计算出的全局 norm 本身是错的——一旦 norm 超过 max_norm 就会按错误的范数过度裁剪。
()
max_norm
bias
grad_norm
观察:replicate_params → 48,ignored_params + 手动 all_reduce → 100+,基线 → 28。
非崩溃类问题,无异常栈,是数值错误:
baseline(all-gather 全量梯度算的范数) : 28 replicate_params 全局 grad_norm : 48 (≈ 在 shard 组上重复计数) ignored_params + 手动 all_reduce : 100+ 首步 loss 不变,仅 grad_norm 偏大
该问题是怎么引起的?
在 FSDP/HSDP(单机 8 卡,DP=8)下,
clip_grad_norm_计算全局梯度范数时,把非切分参数(replicate 参数,或经 FSDPignored_params由用户自管理的参数)的本地 norm² 也放进了 shard 进程组的all_reduce(SUM)。这些梯度本身在各 rank 上已是全局一致的值,其 norm² 应只在本地计入一次;多做一次 shard 组 SUM 会把它重复计shard_world_size次(8 卡即 ×8),使全局 grad_norm 偏大。两种管理不切参数的方式都复现同一症状(并行策略未改、首步 loss 未变,仅 grad_norm 改变):
replicate_params管理不切参数 → global grad_norm = 48(基线 28);ignored_params并在训练侧自行额外 all_reduce 其梯度 → grad_norm = 100+。根因:replicate /
()-签名 的梯度应"本地计入、不通信",旧实现却在 shard 组上对其 norm² 求和。注:首步 loss 不变是因为该步 norm 未超过max_norm、未触发实际裁剪,但计算出的全局 norm 本身是错的——一旦 norm 超过max_norm就会按错误的范数过度裁剪。重现步骤
bias):方式 A 交给replicate_params;方式 B 放入 FSDPignored_params并在训练侧自行对其梯度做一次 all_reduce。clip_grad_norm_,读取返回值/日志中的全局grad_norm。观察:
replicate_params→ 48,ignored_params+ 手动 all_reduce → 100+,基线 → 28。报错信息
非崩溃类问题,无异常栈,是数值错误: