已开启
[基础能力2] HSDP / FSDP 对标 PyTorch:差距梳理与补齐路线图 #5
changzherui创建于  6月30日
changzherui
changzherui成员
6月30日 创建

一、目标与背景

本 Issue 目标: 在保持 Hyper 现有架构不变的前提下,完善 HSDP / FSDP 数据并行能力,使其在语义与 API 上尽可能对齐 PyTorch FSDP2(torch.distributed.fsdp.fully_shard)及 TorchTitan 训练编排习惯,降低 PyTorch 生态迁移成本。

上级文档: 本 Issue 是 HyperParallel / PyTorch / TorchTitan 用户接口对标 在 HSDP/FSDP 模块的子专项,聚焦「参数切分 + 梯度同步 + 训练编排」层的差距分析与补齐计划。

分析依据(源码):

侧 路径
PyTorch torch/distributed/fsdp/_fully_shard/
HyperParallel hyper_parallel/core/fully_shard、platform/torch/fully_shard、platform/mindspore/fully_shard
TorchTitan torchtitan/distributed/fsdp.py

一致性图例(附录中使用): ✅ 一致 · ⚠️ 部分一致 · ❌ 仅一侧有 · 🔶 Hyper 扩展


二、范围与架构约束

2.1 模块范围

本 Issue 覆盖以下子模块的对标与补齐:

子模块 PyTorch Hyper TorchTitan
入口 API _fully_shard/_fully_shard.py core/fully_shard/api.py distributed/fsdp.py
模块状态 _fsdp_state.py、_fsdp_param_group.py core/fully_shard/hsdp_state.py、hsdp_scheduler.py 使用 PyTorch FSDPModule
参数生命周期 _fsdp_param.py platform/*/fully_shard/param.py、param_group.py —
策略与 Mesh _fsdp_api.py、_fsdp_common.py core/fully_shard/utils.py 构造 MixedPrecisionPolicy 等
Hook / 调度 _fsdp_state.py(forward/backward hook) platform/*/fully_shard/scheduler.py、hook_function.py apply_fsdp_to_decoder 编排

2.2 架构约束(补齐前提)

以下设计保持不变,补齐工作在此基础上演进:

  • 公开入口仍为 fully_shard() + 动态混入的 HSDPModule,不改为 PyTorch @contract + MRO FSDPModule 实现
  • 双栈保留:Torch 与 MindSpore 各自 platform/*/fully_shard/,共享 core/fully_shard 调度契约
  • 参数以 local shard(DTensor 或 plain tensor)管理,forward 前 all-gather、backward 后 reduce-scatter 的生命周期模型不变
  • 算子层仍依赖 DTensor 注册式 dispatch;FSDP 不改为 PyTorch 式全量 __torch_dispatch__

2.3 与 PyTorch 的三条结构性差异

理解以下差异,是阅读后文差距与路线图的前提:

  1. 模块扩展方式:PyTorch 通过 MRO 插入 FSDPModule(@contract + FSDPState);Hyper 动态扩展为 HSDP{OrigClass},调度器为 HSDPSchedulerV2。
  2. full_dtensor 集成:PyTorch 提供 DataParallelMeshDims,可从多维 SPMD mesh 显式抽取/flatten DP 子维;Hyper 当前要求用户自行 mesh.flatten() 或传入已 flatten 的一维 mesh(如 CP+FSDP 场景)。
  3. 平台覆盖:PyTorch 仅 Torch;Hyper 同时服务 NPU Torch + MindSpore,部分 API(comm_fusion_zero_copy 默认值、梯度 apply 时机)按后端分化。

2.4 并行模式速览

模式 mesh 维度 参数排布 梯度同步 典型场景
FSDP 1D Shard(0) 全切分 reduce-scatter 最大显存节省
HSDP 2D Replicate() × Shard(0) reduce-scatter + all-reduce 折中通信与显存
复制参数 — 不切分 all-reduce Hyper replicate_params / PyTorch ReplicateModule

三、对标结论(速览)

3.1 已对齐的能力(可直接复用思路)

能力 说明
fully_shard(module, mesh=...) FSDP2 核心入口;1D→FSDP、2D→HSDP 语义一致
fully_shard([m1, m2, ...]) 多模块合并为一个 FSDP unit,共享 collective
reshard_after_forward / set_reshard_after_forward 正向后 reshard 省显存;运行时可切换
unshard / reshard / UnshardHandle 手动与异步 unshard
set_requires_gradient_sync 梯度累积时关闭通信同步
set_requires_all_reduce HSDP 下独立控制 replicate 维 all-reduce
set_modules_to_{forward,backward}_prefetch 层间 prefetch
MixedPrecisionPolicy / CPUOffloadPolicy 混合精度与 CPU offload 策略类
ignored_params 排除在 FSDP 生命周期之外
shard_placement_fn 逐参数自定义切分维(MoE/EP 场景)

3.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 shard
  • apply_grad_on_fp32_main_grad(Torch 路径 FP32 main grad)
  • MindSpore 双栈:enable_mindspore_backward_compat() + 平台分化 post_backward

3.3 按场景一致性总表

用户场景 一致性 说明
单层 / 逐层 fully_shard + 1D mesh ✅ 与 PyTorch FSDP2 用法一致
2D mesh HSDP ✅ replicate 维 0、shard 维 1
fully_shard + TP 组合(先 TP 后 FSDP) ⚠️ 需关注 DTensor mesh 维数一致;部分 mesh 组合仍有限制
CP + FSDP(dp_shard + cp flatten) ⚠️ Hyper 需手动 flatten mesh;PyTorch 用 dp_mesh_dims
MoE + EP 分参 shard_placement_fn ⚠️ Titan/PyTorch 支持 per-param mesh;Hyper 能力较简
梯度累积 no_sync ✅ set_requires_gradient_sync(False)
Titan 全局 token 梯度缩放 ⚠️ Titan disable_fsdp_gradient_division;Hyper 用 set_gradient_scaling_factor
checkpoint 加载到分片参数 ⚠️ Hyper 自定义 load_state_dict;PyTorch 走 DCP
forward/backward 异常后恢复 ❌ PyTorch 有 reset_iter_state;Hyper 无
MindSpore 后端 FSDP 训练 🔶 Hyper 独有
对称内存 / NCCL 通信调优 ❌ PyTorch set_symm_mem_for_comm 等;Hyper 无对等 API

3.4 迁移三原则

  1. 类名替换:FSDPModule → HSDPModule;from hyper_parallel import fully_shard, HSDPModule。
  2. 多维 mesh 勿直接传 dp_mesh_dims——补齐前将 dp_shard + cp 等子维 flatten 为一维 mesh 再传入 fully_shard(见 docs/guide/context_parallel.md)。
  3. 梯度缩放语义自行对齐:Titan 的 set_gradient_divide_factor(1.0) 对应 Hyper 的 set_reduce_op_type + set_gradient_scaling_factor,勿假设默认行为相同。

四、能力差距(PyTorch 可达 · Hyper 未达)

在架构约束下,以下差距需通过功能补齐弥合;❌ 不支持 · ⚠️ 有替代但不等价。

4.1 阻塞级差距(对应 P0)

差距 PyTorch Hyper 影响
DataParallelMeshDims / dp_mesh_dims 从 full_dtensor mesh 抽取 DP 轴 ❌ 需用户手动 flatten CP+TP+FSDP 多维 mesh 编排
reset_iter_state ✅ ❌ forward/backward 异常后无法安全恢复迭代
reshard_after_forward: int 分层策略 ✅ ❌ 仅 bool 按层差异化 reshard(Titan PP 场景)
per-param mesh via ShardPlacementResult ✅ ⚠️ shard_placement_fn 仅返回 Shard MoE+EP 每参数不同 mesh
register_fsdp_forward_method ✅ ❌ 非 forward 方法需手动 unshard

4.2 高级并行差距(对应 P1)

差距 PyTorch Hyper 影响
ReplicateModule / replicate() ✅ composable API ⚠️ 仅 replicate_params 集合 API 风格不一致
share_comm_ctx ✅ 多 unit 共享通信上下文 ❌ 多 FSDP group 显存/流优化
set_custom_all_gather / set_custom_reduce_scatter ✅ ❌ 自定义通信 backend
set_all_reduce_hook ✅ ❌ 自定义 AR 逻辑
set_symm_mem_for_comm / set_force_sum_reduction_for_comms ✅ ❌ 对称内存通信优化
set_post_optim_event ✅ ❌ optimizer 与下一轮 AG stream 协同
set_reduce_scatter_unused_params 等 RS 细粒度控制 ✅ ❌ 未使用参数 / buffer 控制
set_unshard_in_backward ✅ ❌ backward unshard 行为调控
FSDP + DCP 深度集成 DTensor 元数据内置 ⚠️ 独立 distributed_checkpoint 模块 checkpoint 互操作
Torch 异步梯度 staging 与 root callback pre_reduce_scatter_params 等 ⚠️ 实现等价但缺部分恢复 API 调试与异常路径

4.3 完备性与生态差距(对应 P2)

差距 PyTorch Hyper
apply_fsdp_to_decoder 级训练编排 Titan 提供 ❌(用户自行逐层 wrap 或依赖 examples)
comm_fusion 默认开启 内建 param group 融合 ⚠️ 默认 False,需显式开启
动转静 FSDP compile 生态推进中 ❌(README 标注待实现)
get_model_state_dict 与 DCP 格式对齐 PyTorch DCP 标准 ⚠️ 有实现,互操作待验证
Dynamo / torch.compile FSDP pass torch/_inductor/fx_passes/fsdp.py ❌

4.4 明确不在本 Issue 补齐范围

项 原因
改为 PyTorch FSDPModule MRO 实现 需推翻 HSDPModule 动态混入架构
NCCL 对称内存原样移植 昇腾通信栈差异,需单独立项
TorchTitan Trainer 整体移植 属训练框架层,非 FSDP 原语层

五、补齐路线图(P0 → P1 → P2)

原则: 不改 HSDPModule + 双栈 + 注册式 dispatch 架构,在现有 fully_shard 上增量补齐。

5.1 优先级总览

优先级 目标 项数
P0 阻塞多维 mesh 迁移 / 异常恢复 4
P1 高级并行调优与 MoE 编排 6
P2 生态完备、性能默认、训练编排 4

5.2 P0 — 必须补

# 任务 现状 补齐方向 工作量
1 DataParallelMeshDims 对等能力 用户手动 flatten 引入 dp_mesh_dims 或 mesh["dp_shard","cp"].flatten() 辅助 API;文档与 CP 指南对齐 中
2 reset_iter_state 无 根模块迭代状态重置:wait collectives、reshard、清 pending reduce 中
3 per-param mesh shard_placement_fn 仅返回 Shard 对齐 ShardPlacementResult(mesh + placement) 中~大
4 register_fsdp_forward_method 无 为非 forward 方法注册 pre/post hook 小~中

5.3 P1 — 应补

# 任务 现状 补齐方向 工作量
5 share_comm_ctx 无 多 FSDP unit 共享 HSDPParamGroup 通信 buffer/stream 中
6 自定义通信注入 无 set_custom_all_gather / set_custom_reduce_scatter NPU 适配 中
7 ReplicateModule 风格 API 仅 replicate_params 提供 composable replicate() 或文档化映射关系 小
8 set_post_optim_event 无 optimizer step 后与 all-gather stream 同步 小
9 RS 细粒度控制 无 set_reduce_scatter_unused_params 等等价能力 中
10 DCP checkpoint 桥接 独立模块 FSDP 分片 state 与 PyTorch DCP 格式互操作验证 + adapter 中~大

5.4 P2 — 建议补

# 任务 补齐方向 工作量
11 comm_fusion 默认策略 评估热路径默认开启 + zero_copy 后端分化文档 小
12 Decoder 级编排参考实现 对标 Titan apply_fsdp_to_decoder 的 Hyper examples / trainer 封装 中
13 reshard_after_forward: int 支持分层 reshard 策略(PP 场景) 中
14 动转静局部 FSDP 与 README 路线图联动,局部高阶图 大

5.5 里程碑

M1(多维 mesh 可迁移)  P0 #1 #4
M2(训练鲁棒性)        P0 #2 #3
M3(MoE/EP 与通信调优)  P1 #5 #6 #9
M4(生态与编排完备)    P1 #10 + P2 #11~#14

5.6 执行优先级结论

  1. dp_mesh_dims 对等 + reset_iter_state — 不改架构下解锁 CP/TP/FSDP 组合迁移与异常恢复。
  2. per-param mesh + share_comm_ctx — 解锁 MoE+EP 与多 group 性能优化。
  3. P2 编排与动转静 — M3 后并行推进,提升开箱即用体验。

附录 A:逐类逐接口对照

详细技术参考;日常决策优先看第三~五章。

A.1 入口 API:fully_shard

参数 / 行为 PyTorch Hyper 一致性 功能描述
module / list[module] ✅ ✅ ✅ 单模块或合并 FSDP unit
mesh 1D FSDP / 2D HSDP 同左 ✅ 设备拓扑
reshard_after_forward bool | int | None bool ⚠️ 正向后 reshard
shard_placement_fn ShardPlacementFnResult Shard | None ⚠️ 逐参数切分
mp_policy / offload_policy ✅ ✅ ✅ 混合精度与 offload
ignored_params ✅ ✅ ✅ 排除参数
replicate_params — ✅ 🔶 Hyper 声明式复制参数
dp_mesh_dims DataParallelMeshDims ❌ ❌ full_dtensor DP 轴声明
comm_fusion 内建 显式开关,默认 False ⚠️ AG/RS 融合
comm_fusion_zero_copy — Torch 默认 True、MS False 🔶 零拷贝 flat buffer

A.2 模块封装:HSDPModule / FSDPModule

接口 PyTorch Hyper 一致性 功能描述
unshard / reshard FSDPModule HSDPModule ✅ 手动参数状态切换
UnshardHandle.wait ✅ _UnshardHandle ✅ 异步 unshard 等待
set_requires_gradient_sync ✅ ✅ ✅ 梯度累积
set_requires_all_reduce ✅ ✅ ✅ HSDP AR 独立控制
set_reshard_after_{forward,backward} ✅ ✅ ✅ 运行时 reshard 策略
set_is_last_backward ✅ ✅ ✅ microbatch 标记
set_modules_to_{forward,backward}_prefetch ✅ ✅ ✅ prefetch
load_state_dict 标准 + DCP 自定义(绕过 dispatch) ⚠️ 分片 checkpoint 加载
reset_iter_state ✅ ❌ ❌ 异常后迭代重置
set_post_optim_event ✅ ❌ ❌ optimizer 后 stream 同步
set_gradient_divide_factor ✅ — ⚠️ Hyper 用 set_reduce_op_type
set_gradient_scaling_factor — ✅ 🔶 reduce 后梯度缩放
set_reduce_op_type — ✅ 🔶 "avg" / "sum"
set_custom_all_gather 等 ✅ ❌ ❌ 自定义通信
set_symm_mem_for_comm ✅ ❌ ❌ 对称内存优化
ReplicateModule ✅ — ⚠️ Hyper 用 replicate_params

A.3 策略类

类型 PyTorch Hyper 一致性 功能描述
MixedPrecisionPolicy param/reduce/output_dtype、cast_forward_inputs 同左 + apply_grad_on_fp32_main_grad ⚠️ 模块级混合精度
OffloadPolicy / CPUOffloadPolicy ✅ ✅(pin_memory) ✅ CPU offload
FSDPMeshInfo / HSDPMeshInfo 内部 core/fully_shard/utils.py ✅ shard/replicate 进程组元数据

A.4 通信与同步

接口 PyTorch Hyper 一致性 功能描述
hsdp_sync_stream — ✅ 🔶 等待梯度异步通信完成
share_comm_ctx ✅ ❌ ❌ 共享通信上下文
register_fsdp_forward_method ✅ ❌ ❌ 非 forward 方法 hook
comm_fusion + HSDPParamGroup 内建 FSDPParamGroup 显式开启 ⚠️ flat param 融合 collective
get_model_state_dict DCP 生态 api.get_model_state_dict ⚠️ 分片 state dict

A.5 TorchTitan 编排(仅 Titan)

接口 Hyper 说明
apply_fsdp_to_decoder ❌ Decoder 逐层/MoE FSDP 编排
get_fsdp_reshard_after_forward_policy ❌ PP 感知的 reshard 策略
disable_fsdp_gradient_division ⚠️ 对应 Hyper 梯度缩放 API
enable_fsdp_symm_mem ❌ 全局对称内存 FSDP

A.6 运行时生命周期对照

初始化
  fully_shard(module, mesh) → 参数切为 local shard
  → 注册 forward pre/post hook、backward hook(PostBackwardFunction)

Forward
  pre-hook: lazy_init → unshard(all-gather)→ [forward prefetch]
  → 计算(mp_policy 输入/参数 dtype 转换)
  post-hook: 若 reshard_after_forward → shard(释放 unsharded 存储)

Backward
  pre-hook: 若已 reshard → 再次 unshard → [backward prefetch]
  → 本地反传,梯度写在 unsharded_param.grad
  post-hook / root callback:
    - sharded params: reduce-scatter
    - HSDP replicate 维: all-reduce
    - replicate_params: DDP 式 all-reduce
  → reshard,写入 sharded_param.grad
要点 PyTorch FSDP2 Hyper HSDP/FSDP
状态机 FSDPState + FSDPParamGroup HSDPState + HSDPSchedulerV2
梯度累积 set_requires_gradient_sync(False) 同左
异常恢复 reset_iter_state() —
MindSpore — post_backward 内联 reduce;avg = SUM + div
通信融合 内建 comm_fusion=True 显式开启

A.7 与 TorchTitan 的关系

  • TorchTitan 不重新定义 FSDP,直接调用 PyTorch fully_shard;Hyper 用户需 from hyper_parallel import fully_shard。
  • Titan apply_fsdp_to_decoder 封装了:逐 transformer block wrap、weight tying 分组、MoE per-param mesh、reshard_after_forward PP 策略、EP prefetch;Hyper 目前在 examples/torch/llama3/ 提供类似组合示例,尚无对等 trainer 级 API。
  • Titan disable_fsdp_gradient_division 映射为 Hyper set_gradient_scaling_factor / set_reduce_op_type,需按全局 token 数自行配置。


父文档:HyperParallel / PyTorch / TorchTitan 用户接口对标 · 关联章节:DTensor / DeviceMesh 子专项

likedislike
changzheruichangzherui成员
6月30日 修改标题为 “[基础能力2] HSDP / FSDP 对标 PyTorch:差距梳理与补齐路线图”,原标题为“[基础能力] HSDP / FSDP 对标 PyTorch:差距梳理与补齐路线图”
changzherui
changzherui成员
7月1日 评论:

已补充 §2 HSDP/FSDP 全量接口对标总表(接口索引附录):#8 HSDP / FSDP 全量接口对标总表

likedislike
changzherui
changzherui成员
7月1日 评论:

已补充 §2 HSDP/FSDP 全量接口对标总表(接口索引附录):#8 HSDP / FSDP 全量接口对标总表

likedislike