已开启
[Bug] params_dtype=bfloat16 时 FSDP apply_reduced_grad 崩溃:reduced grad(fp32)与 param shard(bf16)dtype 不一致 #215
Java不加糖创建于 6月15日
6月16日 关联了pull request:fix(mindspore): restore HSDP RS/AR overlap and fix bf16 dtype regressions
6月16日 关联了pull request:fix(mindspore): restore HSDP RS/AR overlap and fix bf16 dtype regressions
7月10日 关联了pull request:fix(mindspore): restore HSDP RS/AR overlap and fix dtype/main_grad regressions (r1.0.0)
7月10日 关联了pull request:fix(mindspore): restore HSDP RS/AR overlap and fix dtype/main_grad regressions (r1.0.0)
问题描述
params_dtype = bfloat16(bf16 master 权重)时,动态图 FSDP 反向在归约后写回梯度处崩溃:形状一致,仅 dtype 不一致:reduced grad 是
Float32,而 param shard 是BFloat16。调用栈(全程在 hyper-parallel FSDP 内):
复现环境
15984d1(!821)params_dtype: bfloat16+compute_dtype: bfloat16,tensor_parallel: 2(2 卡即可复现);崩溃 weight 形状{64640,1024}= vocab(129280/TP2)× hidden(1024)根因分析
FSDP
apply_reduced_grad(param.py:920)在把 reduce-scatter 后的梯度写回 param shard 时,经autograd_compat.py:75 grad校验 grad.dtype 必须等于 source(param)dtype。当params_dtype=bf16时:标准混合精度(
params_dtype=float32master +compute_dtype=bfloat16)下 grad 与 param 都是 fp32,故不触发;一旦 master 权重设为 bf16 即崩。建议
FSDP 在
apply_reduced_grad写回前,应把 reduced grad cast 到 param 的 dtype(或在 grad 不变量校验时允许 dtype 不同、由框架统一 cast),而不是直接断言grad.dtype == param.dtype。这样 bf16 master 权重(grad 在 fp32 归约后降回 bf16)才能正常训练。影响
任何
params_dtype=bfloat16(bf16 master 权重)的动态图 FSDP 训练首个 step 即崩,无法使用纯 bf16 master 权重配置。