Dual-mode Trainer 通过顶层 fsdp_config 配置 FSDP/HSDP 行为,并在分布式拓扑形成有效 FSDP domain 时创建 FSDP2Manager。
fsdp_config
FSDP2Manager
本 issue 交付以下能力:
当前 dual-mode Trainer 是 PyTorch 路径,本 issue 不以 MindSpore dual-mode 为验收范围。
FSDP2Manager 不负责重复构造 rank layout。分布式基础设施统一构造 device_mesh、fsdp_non_moe_mesh 和 fsdp_moe_mesh,Manager 负责选择对应 sub-mesh、解析配置、划分 FSDP unit,并调用 fully_shard()。
device_mesh
fsdp_non_moe_mesh
fsdp_moe_mesh
fully_shard()
以下为 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。
enabled
dp_shard_size
edp_shard_size
fsdp_config.dp_shard_size
1
fsdp_config.edp_shard_size
ep_size > 1
fsdp_config.replicate_params
[]
dp_replicate_size 不是 YAML 配置项,由运行时拓扑推导:
dp_replicate_size
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。
mix_precision.param_dtype
null
mix_precision.reduce_dtype
mix_precision.output_dtype
mix_precision.cast_forward_inputs
true
param_dtype
mix_precision.fp32_main_grad
false
main_grad
三个 dtype 支持:
float16 bfloat16 float32
三个 dtype 都为 null 时,不执行显式 FSDP mixed-precision dtype 转换。
enable_offload
reshard_after_forward
reshard_after_backward
requires_grad_sync
root module 的 reshard_after_forward 固定为 False,YAML 中的 reshard_after_forward 作用于 child FSDP unit。
False
reshard_after_backward 和 requires_grad_sync 由 Trainer 训练循环消费,不由 FSDP2Manager.parallelize() 直接消费。
FSDP2Manager.parallelize()
forward_prefetch_depth
0
backward_prefetch_depth
comm_fusion
comm_fusion_zero_copy
comm_fusion=true
Prefetch 顺序依据模型 module traversal/declaration 顺序生成,不按 dense/expert mesh 分组。
FSDP2Manager.parallelize() 执行以下流程:
Partial
replicate_params
reshard_after_forward=False
Manager 当前使用固定 transformer-block wrap 规则,YAML 不暴露自定义 wrap policy。
以下组合不支持,必须在训练开始前给出明确错误:
dp_shard_size < 1
edp_shard_size < 1
forward_prefetch_depth < 0
backward_prefetch_depth < 0
world_size
tp_size × cp_size × pp_size
ep_size
fp32_main_grad=true
reduce_dtype
float32
fsdp_config.mix_precision.fp32_main_grad
optimizer.fp32_main_params
compile.fullgraph=true
以下不是互斥配置,但只有满足前置条件才有可观察效果:
cast_forward_inputs
除上述约束外,不额外定义 offload、mixed precision、prefetch、通信融合之间的配置优先级。
FSDP2Config
至少覆盖:
纯 FSDP shard replicate × shard HSDP
验收结果:
至少覆盖具有代表性的组合:
FSDP + TP FSDP + CP FSDP + TP + CP
分别验证:
output_dtype
cast_forward_inputs=true
cast_forward_inputs=false
完整开启:
fsdp_config: mix_precision: reduce_dtype: float32 fp32_main_grad: true optimizer: fp32_main_params: true
enable_offload=true
reshard_after_forward=true
reshard_after_forward=false
reshard_after_backward=false
在 dp_shard_size > 1 且一个 optimizer step 包含多个 micro-batch 的场景验证:
dp_shard_size > 1
requires_grad_sync=true
requires_grad_sync=false
forward_prefetch_depth=0
backward_prefetch_depth=0
comm_fusion_zero_copy=false
comm_fusion_zero_copy=null
comm_fusion_zero_copy=true
Mesh 构造保留以下内部约束,但不作为端到端测试的主要验收视角:
(dp, cp, tp)
(fsdp_replicate, fsdp_shard, tp)
Mesh shape、sub-mesh 来源和 concat 关系由 UT/ST 做内部结构验证;端到端验收以 YAML 行为、训练数值、dtype、collective trace、显存及错误信息为准。
Dual-mode Trainer:FSDP2Manager 配置化接入
1. 目标与范围
Dual-mode Trainer 通过顶层
fsdp_config配置 FSDP/HSDP 行为,并在分布式拓扑形成有效 FSDP domain 时创建FSDP2Manager。本 issue 交付以下能力:
当前 dual-mode Trainer 是 PyTorch 路径,本 issue 不以 MindSpore dual-mode 为验收范围。
FSDP2Manager不负责重复构造 rank layout。分布式基础设施统一构造device_mesh、fsdp_non_moe_mesh和fsdp_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: falseFSDP 没有单独的
enabled开关。Trainer 根据 world size、parallel topology、dp_shard_size、推导出的 replicate size 和edp_shard_size判断是否需要创建FSDP2Manager。单卡或不存在有效 FSDP domain 时跳过 FSDP wrap。2.1 拓扑与参数配置
fsdp_config.dp_shard_size1fsdp_config.edp_shard_size1ep_size > 1时有意义fsdp_config.replicate_params[]dp_replicate_size不是 YAML 配置项,由运行时拓扑推导:CP 不扩大
dp_shard_size指定的参数 shard degree。2.2 混合精度配置
mix_precision.param_dtypenullmix_precision.reduce_dtypenullmix_precision.output_dtypenullmix_precision.cast_forward_inputstrueparam_dtype时,将浮点 forward 输入转换到参数计算 dtypemix_precision.fp32_main_gradfalsemain_grad,供 fp32 main-param optimizer wrapper 使用三个 dtype 支持:
三个 dtype 都为
null时,不执行显式 FSDP mixed-precision dtype 转换。2.3 内存与参数生命周期
enable_offloadfalsereshard_after_forwardtruereshard_after_backwardtruerequires_grad_synctrueroot module 的
reshard_after_forward固定为False,YAML 中的reshard_after_forward作用于 child FSDP unit。reshard_after_backward和requires_grad_sync由 Trainer 训练循环消费,不由FSDP2Manager.parallelize()直接消费。2.4 Prefetch 与通信融合
forward_prefetch_depth10表示关闭 forward prefetchbackward_prefetch_depth10表示关闭 backward prefetchcomm_fusionfalsecomm_fusion_zero_copynullnull在comm_fusion=true时默认开启,false使用 copy-in 路径Prefetch 顺序依据模型 module traversal/declaration 顺序生成,不按 dense/expert mesh 分组。
3. FSDP2Manager 行为
FSDP2Manager.parallelize()执行以下流程:Partialplacement 或 tied parameter 冲突布局。replicate_params。reshard_after_forward=False。Manager 当前使用固定 transformer-block wrap 规则,YAML 不暴露自定义 wrap policy。
4. 配置约束
4.1 配置阶段必须失败的组合
以下组合不支持,必须在训练开始前给出明确错误:
dp_shard_size < 1。edp_shard_size < 1。forward_prefetch_depth < 0或backward_prefetch_depth < 0。world_size不能被tp_size × cp_size × pp_size整除。dp_shard_size整除。ep_size整除。edp_shard_size整除。fp32_main_grad=true,但reduce_dtype未配置为float32。fsdp_config.mix_precision.fp32_main_grad与optimizer.fp32_main_params没有同时开启或同时关闭。compile.fullgraph=true。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_backward、requires_grad_sync:需要一个 optimizer step 包含多个 micro-batch。replicate_params:必须填写最终模型可以解析的参数 FQN。除上述约束外,不额外定义 offload、mixed precision、prefetch、通信融合之间的配置优先级。
5. 验收点
5.1 YAML 解析与默认值
FSDP2Config的完整默认值。5.2 FSDP 启用、跳过和 wrap
FSDP2Manager。reshard_after_forward控制。5.3 Dense FSDP 与 HSDP
至少覆盖:
验收结果:
dp_shard_size。5.4 FSDP 与 TP、CP、PP 组合
至少覆盖具有代表性的组合:
验收结果:
5.5 MoE EP/EDP
ep_size > 1时构造 routed expert 的独立 EDP FSDP domain。dp_shard_size,expert 参数使用edp_shard_size。5.6
replicate_paramsreplicate_params的数值参考在预期语义下对齐。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 自动转换输入。null时不进行显式 dtype 转换。5.8 fp32 main-grad
完整开启:
fsdp_config: mix_precision: reduce_dtype: float32 fp32_main_grad: true optimizer: fp32_main_params: true验收结果:
main_grad。reduce_dtype或配置非 float32 时必须失败,不能静默修正。5.9 CPU offload 与 reshard
enable_offload=true时,sharded 参数和梯度在非计算阶段位于 CPU。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,最后一个必须同步。5.11 Prefetch 与通信融合
forward_prefetch_depth=0、backward_prefetch_depth=0时不设置 prefetch target。comm_fusion=true时,多参数 FSDP unit 的 all-gather/reduce-scatter collective 次数减少并使用融合 buffer。comm_fusion_zero_copy=false使用 copy-in 路径。comm_fusion=true、comm_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)根域。fsdp_non_moe_mesh。fsdp_moe_mesh保持 EDP shard/replicate 和 EP 根域。fsdp_moe_mesh。Mesh shape、sub-mesh 来源和 concat 关系由 UT/ST 做内部结构验证;端到端验收以 YAML 行为、训练数值、dtype、collective trace、显存及错误信息为准。