已开启
[基础能力2] HSDP / FSDP 对标 PyTorch:差距梳理与补齐路线图 #5
changzherui创建于 6月30日
6月30日 修改标题为 “[基础能力2] HSDP / FSDP 对标 PyTorch:差距梳理与补齐路线图”,原标题为“[基础能力] HSDP / FSDP 对标 PyTorch:差距梳理与补齐路线图”
6月30日 修改标题为 “[基础能力2] HSDP / FSDP 对标 PyTorch:差距梳理与补齐路线图”,原标题为“[基础能力] HSDP / FSDP 对标 PyTorch:差距梳理与补齐路线图”
changzherui
7月1日 评论:
7月1日 评论:
已补充 §2 HSDP/FSDP 全量接口对标总表(接口索引附录):#8 HSDP / FSDP 全量接口对标总表


changzherui
7月1日 评论:
7月1日 评论:
已补充 §2 HSDP/FSDP 全量接口对标总表(接口索引附录):#8 HSDP / FSDP 全量接口对标总表


一、目标与背景
本 Issue 目标: 在保持 Hyper 现有架构不变的前提下,完善 HSDP / FSDP 数据并行能力,使其在语义与 API 上尽可能对齐 PyTorch FSDP2(
torch.distributed.fsdp.fully_shard)及 TorchTitan 训练编排习惯,降低 PyTorch 生态迁移成本。上级文档: 本 Issue 是 HyperParallel / PyTorch / TorchTitan 用户接口对标 在 HSDP/FSDP 模块的子专项,聚焦「参数切分 + 梯度同步 + 训练编排」层的差距分析与补齐计划。
分析依据(源码):
torch/distributed/fsdp/_fully_shard/hyper_parallel/core/fully_shard、platform/torch/fully_shard、platform/mindspore/fully_shardtorchtitan/distributed/fsdp.py一致性图例(附录中使用): ✅ 一致 · ⚠️ 部分一致 · ❌ 仅一侧有 · 🔶 Hyper 扩展
二、范围与架构约束
2.1 模块范围
本 Issue 覆盖以下子模块的对标与补齐:
_fully_shard/_fully_shard.pycore/fully_shard/api.pydistributed/fsdp.py_fsdp_state.py、_fsdp_param_group.pycore/fully_shard/hsdp_state.py、hsdp_scheduler.pyFSDPModule_fsdp_param.pyplatform/*/fully_shard/param.py、param_group.py_fsdp_api.py、_fsdp_common.pycore/fully_shard/utils.pyMixedPrecisionPolicy等_fsdp_state.py(forward/backward hook)platform/*/fully_shard/scheduler.py、hook_function.pyapply_fsdp_to_decoder编排2.2 架构约束(补齐前提)
以下设计保持不变,补齐工作在此基础上演进:
fully_shard()+ 动态混入的HSDPModule,不改为 PyTorch@contract+ MROFSDPModule实现platform/*/fully_shard/,共享core/fully_shard调度契约__torch_dispatch__2.3 与 PyTorch 的三条结构性差异
理解以下差异,是阅读后文差距与路线图的前提:
FSDPModule(@contract+FSDPState);Hyper 动态扩展为HSDP{OrigClass},调度器为HSDPSchedulerV2。DataParallelMeshDims,可从多维 SPMD mesh 显式抽取/flatten DP 子维;Hyper 当前要求用户自行mesh.flatten()或传入已 flatten 的一维 mesh(如 CP+FSDP 场景)。comm_fusion_zero_copy默认值、梯度 apply 时机)按后端分化。2.4 并行模式速览
Shard(0)全切分Replicate()×Shard(0)replicate_params/ PyTorchReplicateModule三、对标结论(速览)
3.1 已对齐的能力(可直接复用思路)
fully_shard(module, mesh=...)fully_shard([m1, m2, ...])reshard_after_forward/set_reshard_after_forwardunshard/reshard/UnshardHandleset_requires_gradient_syncset_requires_all_reduceset_modules_to_{forward,backward}_prefetchMixedPrecisionPolicy/CPUOffloadPolicyignored_paramsshard_placement_fn3.2 Hyper 独有能力(保留,不作为补齐对象)
HSDPModule命名与hsdp_sync_stream()梯度通信同步comm_fusion+comm_fusion_zero_copy显式融合通信路径replicate_params参数集合声明式 DDP 式梯度同步set_gradient_scaling_factor/set_reduce_op_type("avg"\|"sum")梯度缩放与归约类型load_state_dict绕过 DTensor dispatch 直写 local shardapply_grad_on_fp32_main_grad(Torch 路径 FP32 main grad)enable_mindspore_backward_compat()+ 平台分化post_backward3.3 按场景一致性总表
fully_shard+ 1D meshfully_shard+ TP 组合(先 TP 后 FSDP)dp_shard + cpflatten)dp_mesh_dimsshard_placement_fnno_syncset_requires_gradient_sync(False)disable_fsdp_gradient_division;Hyper 用set_gradient_scaling_factorload_state_dict;PyTorch 走 DCPreset_iter_state;Hyper 无set_symm_mem_for_comm等;Hyper 无对等 API3.4 迁移三原则
FSDPModule→HSDPModule;from hyper_parallel import fully_shard, HSDPModule。dp_mesh_dims——补齐前将dp_shard + cp等子维 flatten 为一维 mesh 再传入fully_shard(见docs/guide/context_parallel.md)。set_gradient_divide_factor(1.0)对应 Hyper 的set_reduce_op_type+set_gradient_scaling_factor,勿假设默认行为相同。四、能力差距(PyTorch 可达 · Hyper 未达)
4.1 阻塞级差距(对应 P0)
DataParallelMeshDims/dp_mesh_dimsreset_iter_statereshard_after_forward: int分层策略boolShardPlacementResultshard_placement_fn仅返回Shardregister_fsdp_forward_methodforward方法需手动 unshard4.2 高级并行差距(对应 P1)
ReplicateModule/replicate()replicate_params集合share_comm_ctxset_custom_all_gather/set_custom_reduce_scatterset_all_reduce_hookset_symm_mem_for_comm/set_force_sum_reduction_for_commsset_post_optim_eventset_reduce_scatter_unused_params等 RS 细粒度控制set_unshard_in_backwarddistributed_checkpoint模块pre_reduce_scatter_params等4.3 完备性与生态差距(对应 P2)
apply_fsdp_to_decoder级训练编排comm_fusion默认开启False,需显式开启get_model_state_dict与 DCP 格式对齐torch.compileFSDP passtorch/_inductor/fx_passes/fsdp.py4.4 明确不在本 Issue 补齐范围
FSDPModuleMRO 实现HSDPModule动态混入架构Trainer整体移植五、补齐路线图(P0 → P1 → P2)
5.1 优先级总览
5.2 P0 — 必须补
DataParallelMeshDims对等能力dp_mesh_dims或mesh["dp_shard","cp"].flatten()辅助 API;文档与 CP 指南对齐reset_iter_stateshard_placement_fnShardShardPlacementResult(mesh + placement)register_fsdp_forward_method5.3 P1 — 应补
share_comm_ctxHSDPParamGroup通信 buffer/streamset_custom_all_gather/set_custom_reduce_scatterNPU 适配ReplicateModule风格 APIreplicate_paramsreplicate()或文档化映射关系set_post_optim_eventset_reduce_scatter_unused_params等等价能力5.4 P2 — 建议补
comm_fusion默认策略apply_fsdp_to_decoder的 Hyper examples / trainer 封装reshard_after_forward: int5.5 里程碑
5.6 执行优先级结论
dp_mesh_dims对等 +reset_iter_state— 不改架构下解锁 CP/TP/FSDP 组合迁移与异常恢复。share_comm_ctx— 解锁 MoE+EP 与多 group 性能优化。附录 A:逐类逐接口对照
A.1 入口 API:
fully_shardmodule/list[module]meshreshard_after_forwardbool | int | Noneboolshard_placement_fnShardPlacementFnResultShard | Nonemp_policy/offload_policyignored_paramsreplicate_paramsdp_mesh_dimsDataParallelMeshDimscomm_fusionFalsecomm_fusion_zero_copyTrue、MSFalseA.2 模块封装:
HSDPModule/FSDPModuleunshard/reshardFSDPModuleHSDPModuleUnshardHandle.wait_UnshardHandleset_requires_gradient_syncset_requires_all_reduceset_reshard_after_{forward,backward}set_is_last_backwardset_modules_to_{forward,backward}_prefetchload_state_dictreset_iter_stateset_post_optim_eventset_gradient_divide_factorset_reduce_op_typeset_gradient_scaling_factorset_reduce_op_type"avg"/"sum"set_custom_all_gather等set_symm_mem_for_commReplicateModulereplicate_paramsA.3 策略类
MixedPrecisionPolicyparam/reduce/output_dtype、cast_forward_inputsapply_grad_on_fp32_main_gradOffloadPolicy/CPUOffloadPolicypin_memory)FSDPMeshInfo/HSDPMeshInfocore/fully_shard/utils.pyA.4 通信与同步
hsdp_sync_streamshare_comm_ctxregister_fsdp_forward_methodcomm_fusion+HSDPParamGroupFSDPParamGroupget_model_state_dictapi.get_model_state_dictA.5 TorchTitan 编排(仅 Titan)
apply_fsdp_to_decoderget_fsdp_reshard_after_forward_policydisable_fsdp_gradient_divisionenable_fsdp_symm_memA.6 运行时生命周期对照
FSDPState+FSDPParamGroupHSDPState+HSDPSchedulerV2set_requires_gradient_sync(False)reset_iter_state()post_backward内联 reduce;avg= SUM + divcomm_fusion=True显式开启A.7 与 TorchTitan 的关系
fully_shard;Hyper 用户需from hyper_parallel import fully_shard。apply_fsdp_to_decoder封装了:逐 transformer block wrap、weight tying 分组、MoE per-param mesh、reshard_after_forwardPP 策略、EP prefetch;Hyper 目前在examples/torch/llama3/提供类似组合示例,尚无对等 trainer 级 API。disable_fsdp_gradient_division映射为 Hyperset_gradient_scaling_factor/set_reduce_op_type,需按全局 token 数自行配置。父文档:HyperParallel / PyTorch / TorchTitan 用户接口对标 · 关联章节:DTensor / DeviceMesh 子专项