已开启
[RFC]: dualmode-trainer 支持FSDPManager #322
MengXY107创建于  8月7日
MengXY107
MengXY107成员
8月7日 创建

Dual-mode Trainer:FSDP2Manager 配置化接入

1. 目标与范围

Dual-mode Trainer 通过顶层 fsdp_config 配置 FSDP/HSDP 行为,并在分布式拓扑形成有效 FSDP domain 时创建 FSDP2Manager

本 issue 交付以下能力:

  • dense/non-MoE transformer block 和 root module 的嵌套 FSDP;
  • routed expert 使用独立 EDP shard domain;
  • FSDP 与 TP、CP、EP、PP 拓扑组合;
  • TP/EP 参数 source layout 接入 FSDP;
  • 参数复制、混合精度、CPU offload、reshard、梯度同步、prefetch 和通信融合;
  • Trainer 梯度累积流程对 FSDP reshard、gradient sync 的控制。

当前 dual-mode Trainer 是 PyTorch 路径,本 issue 不以 MindSpore dual-mode 为验收范围。

FSDP2Manager 不负责重复构造 rank layout。分布式基础设施统一构造 device_meshfsdp_non_moe_meshfsdp_moe_mesh,Manager 负责选择对应 sub-mesh、解析配置、划分 FSDP unit,并调用 fully_shard()

2. YAML 配置接口

以下为 FSDP 相关配置片段:

accelerator:
  tp_size: 1
  cp_size: 1
  ep_size: 1
  pp_size: 1
  sequence_parallel: false
  loss_parallel: false

fsdp_config:
  dp_shard_size: 4
  edp_shard_size: 1
  replicate_params: []

  mix_precision:
    param_dtype: bfloat16
    reduce_dtype: float32
    output_dtype: bfloat16
    cast_forward_inputs: true
    fp32_main_grad: false

  enable_offload: false
  reshard_after_forward: true
  reshard_after_backward: true
  requires_grad_sync: true

  forward_prefetch_depth: 1
  backward_prefetch_depth: 1
  comm_fusion: false
  comm_fusion_zero_copy: null

optimizer:
  _target_: <optimizer target>
  fp32_main_params: false

FSDP 没有单独的 enabled 开关。Trainer 根据 world size、parallel topology、dp_shard_size、推导出的 replicate size 和 edp_shard_size 判断是否需要创建 FSDP2Manager。单卡或不存在有效 FSDP domain 时跳过 FSDP wrap。

2.1 拓扑与参数配置

YAML 配置 默认值 效果
fsdp_config.dp_shard_size 1 dense/non-MoE 参数的 FSDP shard degree;剩余 DP×CP domain 自动成为 replicate domain
fsdp_config.edp_shard_size 1 routed expert 参数的 EDP shard degree,仅在 ep_size > 1 时有意义
fsdp_config.replicate_params [] 按最终模型参数 FQN 指定不分片参数;参数保持 replicated,但梯度仍在所属 FSDP domain 做 all-reduce

dp_replicate_size 不是 YAML 配置项,由运行时拓扑推导:

dp_size = world_size / (tp_size × cp_size × pp_size)
fsdp_data_parallel_size = dp_size × cp_size
dp_replicate_size = fsdp_data_parallel_size / dp_shard_size

CP 不扩大 dp_shard_size 指定的参数 shard degree。

2.2 混合精度配置

YAML 配置 默认值 效果
mix_precision.param_dtype null unshard 后参数参与 forward/backward 计算的 dtype
mix_precision.reduce_dtype null reduce-scatter/all-reduce 使用的梯度通信 dtype
mix_precision.output_dtype null FSDP module 输出转换到的 dtype
mix_precision.cast_forward_inputs true 设置了 param_dtype 时,将浮点 forward 输入转换到参数计算 dtype
mix_precision.fp32_main_grad false FSDP 将规约后的梯度写入 fp32 main_grad,供 fp32 main-param optimizer wrapper 使用

三个 dtype 支持:

float16
bfloat16
float32

三个 dtype 都为 null 时,不执行显式 FSDP mixed-precision dtype 转换。

2.3 内存与参数生命周期

YAML 配置 默认值 效果
enable_offload false 使用 CPU offload;sharded 参数和梯度在非计算阶段保存在 CPU,计算前后进行 H2D/D2H
reshard_after_forward true child FSDP unit forward 后立即 reshard,降低峰值显存;关闭后保留 unsharded 参数供 backward 使用
reshard_after_backward true 控制梯度累积期间 backward 后是否立即 reshard;关闭时非最后一个 micro-batch 可继续保留 unsharded 参数,最后一个 micro-batch 强制恢复 reshard
requires_grad_sync true 多 micro-batch 梯度累积时是否每个 micro-batch 都执行 FSDP gradient sync;关闭时非最后一个 micro-batch 跳过,最后一个必须同步

root module 的 reshard_after_forward 固定为 False,YAML 中的 reshard_after_forward 作用于 child FSDP unit。

reshard_after_backwardrequires_grad_sync 由 Trainer 训练循环消费,不由 FSDP2Manager.parallelize() 直接消费。

2.4 Prefetch 与通信融合

YAML 配置 默认值 效果
forward_prefetch_depth 1 为当前 FSDP unit 预取后续 N 个 unit;0 表示关闭 forward prefetch
backward_prefetch_depth 1 backward 时预取前序 N 个 unit;0 表示关闭 backward prefetch
comm_fusion false 融合 FSDP unit 内的 all-gather 和 reduce-scatter,减少 collective 次数
comm_fusion_zero_copy null 控制通信融合的零拷贝存储路径;PyTorch 下 nullcomm_fusion=true 时默认开启,false 使用 copy-in 路径

Prefetch 顺序依据模型 module traversal/declaration 顺序生成,不按 dense/expert mesh 分组。

3. FSDP2Manager 行为

FSDP2Manager.parallelize() 执行以下流程:

  1. 将 planner 提供的参数 FQN source layout 解析到模型最终 Parameter。
  2. 校验 TP/EP source metadata,不接受未知 FQN、Partial placement 或 tied parameter 冲突布局。
  3. 根据最终参数 FQN 解析 replicate_params
  4. 找到 transformer block,并补充 routed expert 的嵌套 FSDP unit。
  5. 参数归属到最深的 FSDP unit,按 bottom-up 顺序 wrap child unit。
  6. dense unit 使用 dense FSDP domain,expert unit 使用 EDP FSDP domain。
  7. 最后 wrap root module,并固定 root reshard_after_forward=False
  8. 对带 TP/EP source layout 的 unit 配置梯度规约和全局平均缩放。
  9. 按 module traversal/declaration 顺序配置 forward/backward prefetch。

Manager 当前使用固定 transformer-block wrap 规则,YAML 不暴露自定义 wrap policy。

4. 配置约束

4.1 配置阶段必须失败的组合

以下组合不支持,必须在训练开始前给出明确错误:

  1. dp_shard_size < 1
  2. edp_shard_size < 1
  3. forward_prefetch_depth < 0backward_prefetch_depth < 0
  4. world_size 不能被 tp_size × cp_size × pp_size 整除。
  5. DP×CP domain 不能被 dp_shard_size 整除。
  6. expert domain 不能被 ep_size 整除。
  7. EDP size 不能被 edp_shard_size 整除。
  8. fp32_main_grad=true,但 reduce_dtype 未配置为 float32
  9. fsdp_config.mix_precision.fp32_main_gradoptimizer.fp32_main_params 没有同时开启或同时关闭。
  10. FSDP 开启时配置顶层 compile.fullgraph=true
  11. replicate_params 或 source metadata 包含模型中不存在的参数 FQN。

4.2 有前置条件才生效的配置

以下不是互斥配置,但只有满足前置条件才有可观察效果:

  • edp_shard_size:需要 ep_size > 1 且模型包含 routed expert。
  • cast_forward_inputs:需要配置非空 param_dtype
  • comm_fusion_zero_copy:需要 comm_fusion=true
  • reshard_after_backwardrequires_grad_sync:需要一个 optimizer step 包含多个 micro-batch。
  • replicate_params:必须填写最终模型可以解析的参数 FQN。

除上述约束外,不额外定义 offload、mixed precision、prefetch、通信融合之间的配置优先级。

5. 验收点

5.1 YAML 解析与默认值

  • 最小 YAML 不填写可选字段时,resolved config 中包含 FSDP2Config 的完整默认值。
  • 配置支持命令行 dotted override。
  • 已删除字段和未知字段必须在配置解析阶段失败,不能静默忽略。
  • 所有非法拓扑、非法 prefetch depth 和非法精度组合返回包含具体配置路径及修复方向的错误。

5.2 FSDP 启用、跳过和 wrap

  • 有效 FSDP/HSDP 拓扑能够完成初始化、至少两个 forward/backward/optimizer step,loss 为有限值。
  • 单卡或不存在有效 FSDP domain 时跳过 FSDP2Manager
  • 支持模型的 transformer block 被逐层 wrap,root module 最后 wrap。
  • 日志可观察被 wrap 的 transformer block 数量。
  • root module forward 后保持 unsharded,child unit 行为由 reshard_after_forward 控制。

5.3 Dense FSDP 与 HSDP

至少覆盖:

纯 FSDP shard
replicate × shard HSDP

验收结果:

  • dense 参数实际 shard degree 等于 dp_shard_size
  • HSDP replicate group 内参数和更新后权重一致。
  • 使用相同 seed 和 batch 时,loss、梯度及更新后权重与单卡或非 FSDP 参考在对应 dtype 容差内对齐。
  • CP 打开后不扩大 dense 参数 shard degree。

5.4 FSDP 与 TP、CP、PP 组合

至少覆盖具有代表性的组合:

FSDP + TP
FSDP + CP
FSDP + TP + CP

验收结果:

  • 各组合可以完成训练 step,不出现 source layout、mesh concat 或 collective group 错误。
  • TP/CP 参数布局在 FSDP wrap 后保持正确。
  • 各 rank loss、梯度同步及更新后权重符合对应并行语义。
  • PP 场景每个 pipeline stage 仅使用当前 stage ranks 构造 FSDP domain。
  • 不以直接断言内部 Mesh 对象身份作为主要验收方式。

5.5 MoE EP/EDP

  • ep_size > 1 时构造 routed expert 的独立 EDP FSDP domain。
  • dense 参数使用 dp_shard_size,expert 参数使用 edp_shard_size
  • expert FSDP unit 使用 EP source metadata,dense unit 使用 TP source metadata。
  • routed expert 不得残留在 root FSDP unit。
  • 相同输入下,loss、expert 梯度和更新后权重与参考实现对齐。
  • 非法 expert/EDP 拓扑在训练开始前失败。

5.6 replicate_params

  • YAML 指定的参数保持完整形状,不参与参数 shard。
  • replicated 参数梯度仍执行 all-reduce,各 rank 更新结果一致。
  • 同一 FSDP unit 可以同时包含 sharded 参数和 replicated 参数。
  • 与不使用 replicate_params 的数值参考在预期语义下对齐。
  • 不存在的参数 FQN 明确报错。

5.7 混合精度

分别验证:

  • param_dtype 控制 forward/backward 参数计算 dtype。
  • reduce_dtype 控制 reduce-scatter/all-reduce dtype。
  • output_dtype 控制 FSDP unit 输出 dtype。
  • cast_forward_inputs=true 时浮点输入转换到 param_dtype
  • cast_forward_inputs=false 时不由 FSDP 自动转换输入。
  • dtype 全为 null 时不进行显式 dtype 转换。
  • float16、bfloat16、float32 配置均能正确解析。
  • 混合精度训练 loss 为有限值,并与对应精度参考在合理容差内对齐。

5.8 fp32 main-grad

完整开启:

fsdp_config:
  mix_precision:
    reduce_dtype: float32
    fp32_main_grad: true

optimizer:
  fp32_main_params: true

验收结果:

  • FSDP 规约后的梯度保存在 fp32 main_grad
  • AdamW/Muon 只更新 fp32 main-param,更新后回刷模型计算参数。
  • 模型计算参数使用低精度时,main-param 和 main-grad 仍为 float32。
  • 只开启任意一侧、未配置 reduce_dtype 或配置非 float32 时必须失败,不能静默修正。

5.9 CPU offload 与 reshard

  • enable_offload=true 时,sharded 参数和梯度在非计算阶段位于 CPU。
  • profiler/trace 中可观察 H2D、D2H,训练可以正常完成。
  • 相同模型下,offload 路径的设备峰值参数内存低于关闭路径。
  • reshard_after_forward=true 时 child unit forward 后释放 unsharded 参数,backward 前重新 all-gather。
  • reshard_after_forward=false 时 child unit 参数保留到 backward,collective 次数或时序发生对应变化。
  • reshard_after_backward=false 且存在多个 micro-batch 时,非最后一个 micro-batch 保留 unsharded 参数,最后一个 micro-batch 恢复 reshard。
  • 各路径最终梯度和更新后权重与参考对齐。

5.10 梯度累积与 gradient sync

dp_shard_size > 1 且一个 optimizer step 包含多个 micro-batch 的场景验证:

  • requires_grad_sync=true 时每个 micro-batch 执行 FSDP gradient sync。
  • requires_grad_sync=false 时非最后一个 micro-batch 不执行 reduce-scatter,最后一个必须同步。
  • HSDP replicate all-reduce 在 optimizer step 最后一个 micro-batch 执行。
  • 两种策略累积出的最终梯度和参数更新结果一致。

5.11 Prefetch 与通信融合

  • forward_prefetch_depth=0backward_prefetch_depth=0 时不设置 prefetch target。
  • depth 为 N 时,trace 中可以观察相邻 N 个 unit 的提前 all-gather。
  • prefetch 不改变 loss、梯度和参数更新结果。
  • comm_fusion=true 时,多参数 FSDP unit 的 all-gather/reduce-scatter collective 次数减少并使用融合 buffer。
  • comm_fusion_zero_copy=false 使用 copy-in 路径。
  • PyTorch 下 comm_fusion=truecomm_fusion_zero_copy=null 使用默认零拷贝路径。
  • 显式 comm_fusion_zero_copy=true 时 optimizer step 能正确更新 view-backed parameter storage。
  • 通信融合及零拷贝开关不改变训练数值。

6. 实现约束

Mesh 构造保留以下内部约束,但不作为端到端测试的主要验收视角:

  • device_mesh 保持 (dp, cp, tp) 拓扑,供 Planner、TP/CP、batch 和 loss 使用。
  • fsdp_non_moe_mesh 保持 (fsdp_replicate, fsdp_shard, tp) 根域。
  • dense FSDP child 和 TP source child 来自同一个 fsdp_non_moe_mesh
  • fsdp_moe_mesh 保持 EDP shard/replicate 和 EP 根域。
  • expert FSDP child 和 EP source child 来自同一个 fsdp_moe_mesh
  • FSDPManager 只选择已经构造的 sub-mesh,不重新计算 rank layout。
  • TP source rank group/local rank 必须与原 TP 拓扑一致。
  • CP 不写入 TP source placement,也不扩大 dense 参数 shard degree。

Mesh shape、sub-mesh 来源和 concat 关系由 UT/ST 做内部结构验证;端到端验收以 YAML 行为、训练数值、dtype、collective trace、显存及错误信息为准。

likedislike
MengXY107MengXY107成员
8月7日 修改了issue 的描述
MengXY107MengXY107成员
26 天前 修改了issue 的描述
MengXY107MengXY107成员
26 天前 修改了issue 的描述
MengXY107MengXY107成员
25 天前 关联了pull request:feat(trainer): 支持 fp32_main_params 与模型精度转换