已开启
[接口梳理2] HSDP / FSDP 全量接口对标总表(Hyper · PyTorch · TorchTitan) #8
changzherui创建于  7月1日
changzherui
changzherui成员
7月1日 创建

一、文档说明

📎 上级文档: HyperParallel / PyTorch / TorchTitan 用户接口对标 #1

📎 姊妹专项(差距与路线图): HSDP / FSDP 对标 PyTorch #5

本 Issue 定位: 在 #1 模块划分下,汇总 §2 HSDP / FSDP 子模块的全量对外接口对标总表(表头:接口名 · 参数 · 功能 · hyper · pytorch · titan),作为 #5 差距梳理的接口索引附录。

分析依据(源码,2026-07):

侧 路径
HyperParallel hyper_parallel/__init__.py、hyper_parallel/core/fully_shard/、platform/*/fully_shard/
PyTorch torch/distributed/fsdp/_fully_shard/、torch/distributed/_composable/replicate_with_fsdp.py
TorchTitan torchtitan/distributed/fsdp.py(编排层,底层调用 PyTorch fully_shard)

图例: ✅ 有对标 · ⚠️ 部分对标 · ❌ 无 · 🔶 Hyper 扩展

列说明:

  • hyper:HyperParallel 导出路径或实现位置
  • pytorch:PyTorch FSDP2 对标 API(torch.distributed.fsdp)
  • titan:TorchTitan 实际调用的接口(多为 PyTorch 原生封装,非 Titan 自有原语)

架构备注(影响对标方式): Hyper 通过动态混入 HSDP{OrigClass} 扩展模块,非 PyTorch @contract + MRO FSDPModule;类名对标为 HSDPModule ↔ FSDPModule。


二、全量接口对标总表

接口名 参数 功能 hyper pytorch titan
入口 API
fully_shard module / list[module] 对模块(或模块列表合并为一个 FSDP unit)应用参数切分与梯度同步 hyper_parallel.fully_shard torch.distributed.fsdp.fully_shard 经 apply_fsdp_to_decoder 逐层调用
fully_shard · mesh DeviceMesh | None 1D→FSDP、2D→HSDP 设备拓扑 ✅ ✅ ✅ dp_mesh / edp_mesh
fully_shard · reshard_after_forward bool(Hyper)/ bool | int | None(PT) 正向后是否 reshard 释放 unsharded 存储 ✅ 仅 bool ✅ bool/int/None get_fsdp_reshard_after_forward_policy 解析为 bool
fully_shard · shard_placement_fn Callable[[Parameter], Shard | None](Hyper) 逐参数自定义切分维 ✅ 返回 Shard | None ✅ 返回 Shard | ShardPlacementResult | None MoE 场景 per-param mesh
fully_shard · mp_policy MixedPrecisionPolicy 模块级混合精度(param/reduce/output dtype、cast 输入) ✅ + apply_grad_on_fp32_main_grad 🔶 ✅ 构造 MixedPrecisionPolicy 传入
fully_shard · offload_policy OffloadPolicy / CPUOffloadPolicy CPU offload 策略 ✅ ✅ cpu_offload=True → CPUOffloadPolicy()
fully_shard · ignored_params set[Parameter] 完全排除在 FSDP 生命周期外 ✅ ✅ —
fully_shard · replicate_params set[Parameter] 参数不切分但参与 DDP 式 all-reduce ✅ 🔶 ❌(用 replicate() / ReplicateModule) —
fully_shard · comm_fusion bool,默认 False 显式开启 AG/RS 融合通信 ✅ 🔶 默认关 内建于 FSDPParamGroup —
fully_shard · comm_fusion_zero_copy Optional[bool] 融合通信零拷贝 flat buffer 路径 ✅ 🔶 PT 默认 True、MS 默认 False ❌ —
fully_shard · dp_mesh_dims DataParallelMeshDims(shard=..., replicate=...) 从 full_dtensor 多维 mesh 抽取/flatten DP 轴 ❌ ✅ ✅ CP+TP+FSDP 组合
fully_shard · 返回值 — 返回原模块类型(动态混入后 isinstance(m, HSDPModule)) HSDP{OrigClass} FSDPModule MRO PyTorch FSDPModule
register_fsdp_forward_method module, method_name 为非 forward 方法注册 pre/post FSDP hook ❌ torch.distributed.fsdp.register_fsdp_forward_method —
share_comm_ctx list[FSDPModule] 多 FSDP unit 共享通信 buffer/stream ❌ torch.distributed.fsdp.share_comm_ctx —
get_cls_to_fsdp_cls — 原类 → FSDP 包装类映射表 ❌ ✅ 内部/调试 —
disable_fsdp_module_new_init 上下文管理器 禁用 FSDPModule.__new__ 特殊构造 ❌ ✅ —
replicate / ReplicateModule module, mesh, … composable DDP 式复制参数 + 梯度 AR ❌ torch.distributed._composable.replicate disable_fsdp_gradient_division 注释兼容
模块封装:HSDPModule / FSDPModule
unshard async_op: bool = False all-gather 展开分片参数 HSDPModule.unshard FSDPModule.unshard 隐式(forward pre-hook)
reshard — 释放 unsharded、恢复 sharded 视图 ✅ ✅ 隐式(forward post-hook)
_UnshardHandle / UnshardHandle — 异步 unshard 句柄 _UnshardHandle.wait() UnshardHandle.wait() —
set_requires_gradient_sync requires_gradient_sync, recurse=True 梯度累积 no_sync(控制 RS+AR) ✅ ✅ —
set_requires_all_reduce requires_all_reduce, recurse=True HSDP 下独立控制 replicate 维 AR ✅ ✅ —
set_reshard_after_forward reshard_after_forward, recurse=True 运行时切换正向后 reshard ✅(recurse=False 未实现) ✅ PP 场景差异化策略
set_reshard_after_backward reshard_after_backward, recurse=True 反向后 reshard(梯度累积换通信) ✅(recurse=False 未实现) ✅ —
set_is_last_backward is_last_backward: bool microbatch 末步标记 ✅ ✅ —
set_modules_to_forward_prefetch modules: list 指定 forward 中显式 prefetch 的模块 ✅ ✅ EP 场景显式设置
set_modules_to_backward_prefetch modules: list 指定 backward 中显式 prefetch 的模块 ✅ ✅ EP 场景显式设置
reset_iter_state — forward/backward 异常后重置迭代状态 ❌ ✅ 根模块调用 —
set_post_optim_event event: torch.Event optimizer 后与 AG stream 同步 ❌ ✅ —
set_gradient_divide_factor factor: float 梯度归约前除法因子(PreMulSum) ❌ ✅ disable_fsdp_gradient_division → 1.0
set_gradient_scaling_factor factor: None | float | Tensor reduce 后梯度乘法缩放 ✅ 🔶 ❌ ⚠️ 映射 Titan 全局 token 缩放
set_reduce_op_type "avg" | "sum" 梯度归约类型 ✅ 🔶 近似 gradient_divide_factor —
set_custom_all_gather comm: AllGather 自定义 all-gather 通信 ❌ ✅ —
set_custom_reduce_scatter comm: ReduceScatter 自定义 reduce-scatter 通信 ❌ ✅ —
set_all_reduce_hook hook, stream=None 自定义 all-reduce 逻辑 ❌ ✅ —
set_force_sum_reduction_for_comms enable: bool 通信强制 SUM 归约(零拷贝前置) ❌ ✅ enable_fsdp_symm_mem 内调用
set_symm_mem_for_comm backend="NCCL" 对称内存 AG 优化 ❌ ✅ enable_fsdp_symm_mem
set_allocate_memory_from_process_group_for_comm enable: bool 从 PG 分配通信 staging buffer ❌ ✅ —
set_reduce_scatter_unused_params reduce_scatter_unused_params, recurse 未使用参数零梯度参与 RS ❌ ✅ MoE 条件分支
set_reduce_scatter_max_input_buffers max_input_buffers, recurse RS copy-in buffer 流水线深度 ❌ ✅ —
set_separate_reduce_scatter_group enable, recurse RS 独立 PG 与 AG 重叠 ❌ ✅ —
set_unshard_in_backward unshard_in_backward: bool 控制 backward 是否 unshard ❌ ✅ embedding 等特例
load_state_dict state_dict, strict=True, assign=False 加载 checkpoint 到分片参数 ✅ 自定义(绕过 DTensor dispatch) 标准 nn.Module + DCP DCP 生态
zero_grad — 清零梯度(MS 路径走 scheduler) ✅ 🔶 继承 nn.Module —
策略与 Mesh 元数据
MixedPrecisionPolicy param_dtype, reduce_dtype, output_dtype, cast_forward_inputs 混合精度策略 dataclass ✅ + apply_grad_on_fp32_main_grad ✅ ✅
OffloadPolicy — 无 offload 默认策略 ✅ ✅ —
CPUOffloadPolicy pin_memory=True CPU offload 参数/梯度 ✅ ✅ cpu_offload=True
DataParallelMeshDims shard, replicate(mesh 轴名) 声明 full_dtensor mesh 中 DP 轴 ❌(内部类字段不同) ✅ 公开 ✅ dp_mesh_dims / edp_mesh_dims
FSDPMeshInfo / HSDPMeshInfo / DDPMeshInfo mesh, shard_mesh_dim, replicate_mesh_dim shard/replicate 进程组元数据 core/fully_shard/utils.py(内部) _fsdp_common.py(内部) 经 _get_mesh_info
ShardPlacementResult placement, mesh_info per-param mesh + placement ❌ _fsdp_common.ShardPlacementResult MoE+EP per-param
AllGather / ReduceScatter / Comm ABC 接口 自定义通信 primitive 契约 ❌ _fsdp_api.py —
通信同步与 State Dict
hsdp_sync_stream — 等待梯度异步通信完成 hyper_parallel.hsdp_sync_stream ❌ —
get_model_state_dict model, options: StateDictOptions 可配置 gather/offload 的 state dict core/fully_shard/api.py(未进 __all__) torch.distributed.checkpoint.state_dict.get_model_state_dict DCP 训练 checkpoint
StateDictOptions full_state_dict, cpu_offload, ignore_frozen_params, broadcast_from_rank0 state dict 导出选项 委托 PyTorch StateDictOptions ✅ —
内部辅助(未公开但 FSDP 链路相关)
get_hsdp_state module 取模块 HSDP 状态 hsdp_utils.get_hsdp_state fully_shard.state(module) —
is_dtensor_managed_param param 判断是否 DTensor 托管参数 hsdp_utils 内部 DTensor 检查 —
infer_fully_shard_param_mode — 推断参数 FSDP/DDP/复制模式 hsdp_utils 内部 —
apply_gradient_scaling_factor — reduce 后梯度缩放实现 hsdp_utils — —
TorchTitan 编排(仅 Titan)
apply_fsdp_to_decoder model, dp_mesh, param_dtype, reduce_dtype, pp_enabled, cpu_offload, reshard_after_forward_policy, ep_degree, edp_mesh, dp_mesh_dims, edp_mesh_dims, enable_symm_mem Decoder 逐层/MoE FSDP 一站式编排 ❌(examples/torch/llama3/ 参考) — ✅ torchtitan/distributed/fsdp.py
get_fsdp_reshard_after_forward_policy policy: str, pp_enabled "always"/"never"/"default" → bool ❌ — ✅
disable_fsdp_gradient_division model 全模型 set_gradient_divide_factor(1.0) ⚠️ 用 set_gradient_scaling_factor 遍历 FSDPModule ✅
enable_fsdp_symm_mem model 全模型对称内存通信优化 ❌ 遍历 FSDPModule ✅

三、差距速查(与 #5 联动)

本节从 总表 中筛出 hyper 列为 ❌ 或 ⚠️ 的项,按来源拆为四块。详细 P0/P1 补齐路线见 #5。

3.1 仅 PyTorch 有 · Hyper 无

接口 / 参数 PyTorch 路径 功能 影响
fully_shard · dp_mesh_dims DataParallelMeshDims 从多维 SPMD mesh 抽取/flatten DP 子轴 CP+TP+FSDP 组合需用户手动 mesh.flatten()
reshard_after_forward: int | None fully_shard 按层/策略差异化 reshard(PP 尾层等) Titan PP 场景 reshard_after_forward_policy 部分语义
shard_placement_fn → ShardPlacementResult _fsdp_common per-param 不同 mesh + placement MoE+EP 每专家不同 FSDP mesh
register_fsdp_forward_method torch.distributed.fsdp 非 forward 方法 FSDP hook 多模态/自定义入口需手动 unshard
share_comm_ctx 同上 多 unit 共享通信上下文 多 FSDP group 显存/流优化
reset_iter_state FSDPModule 异常后 wait collective、reshard、清 pending 训练鲁棒性;失败 microbatch 序列作废
set_post_optim_event FSDPModule optimizer 与下一轮 AG stream 协同 避免 false stream 依赖
set_custom_all_gather / set_custom_reduce_scatter FSDPModule + Comm ABC 注入自定义通信 NPU/异构 backend 调优
set_all_reduce_hook FSDPModule 自定义 HSDP all-reduce 特殊梯度聚合逻辑
set_symm_mem_for_comm / set_force_sum_reduction_for_comms FSDPModule NCCL 对称内存 / SUM-only 通信 多节点 AG 性能(昇腾需单独立项)
set_allocate_memory_from_process_group_for_comm FSDPModule PG 优化 allocator SHARP / 零拷贝
set_reduce_scatter_unused_params FSDPModule 条件分支未用参数零梯度 RS MoE 等多路径模型
set_reduce_scatter_max_input_buffers FSDPModule RS buffer 流水线 暴露 RS 时计算 stall
set_separate_reduce_scatter_group FSDPModule RS 独立 PG AG/RS 网络级重叠
set_unshard_in_backward FSDPModule 跳过 backward unshard embedding 等不参与反传参数
replicate / ReplicateModule _composable.replicate composable 复制参数 API API 风格;Hyper 用 replicate_params 集合
get_cls_to_fsdp_cls / disable_fsdp_module_new_init _fully_shard FSDP 包装类与构造控制 容器索引等边缘场景

3.2 仅 TorchTitan 有 · Hyper 无

Titan 不重新定义 FSDP 原语,底层均为 PyTorch fully_shard + FSDPModule。下表为 Titan 训练编排层 特有、Hyper 无对等公开 API 的项。

接口 / 能力 Titan 路径 功能 影响
apply_fsdp_to_decoder torchtitan/distributed/fsdp.py tok_embeddings/norm/lm_head/逐 block/root 分层 wrap、weight tying 分组、MoE per-param mesh、EP prefetch 开箱即用 Decoder FSDP;Hyper 需 examples 手写
get_fsdp_reshard_after_forward_policy 同上 PP 感知 reshard 策略字符串解析 default = PP 时默认不 reshard
disable_fsdp_gradient_division 同上 遍历设置 gradient_divide_factor=1.0 全局 token 梯度缩放由训练 loop 负责
enable_fsdp_symm_mem 同上 遍历启用对称内存 FSDP 多节点性能调优一键开关
apply_fsdp_to_decoder 内 EP prefetch 链 同上 tok_embeddings→block→…→lm_head 显式 forward/backward prefetch EP 下 D2H 干扰隐式 prefetch 时的补偿

3.3 PyTorch 有 · Hyper 有但语义偏弱(⚠️)

接口 / 参数 PyTorch Hyper 现状 差距说明
reshard_after_forward bool | int | None 仅 bool 无法表达分层 reshard;Titan PP 尾层优化需运行时 set_reshard_after_forward 变通
shard_placement_fn 返回 ShardPlacementResult(含 mesh_info) 返回 Shard | None MoE 专家与 dense 不同 mesh 需补齐
set_reshard_after_{forward,backward} · recurse=False 支持单模块 NotImplementedError 仅 recurse=True 等价实现
set_requires_gradient_sync · recurse 可选 Hyper 固定递归子模块 细粒度控制受限
comm_fusion 内建默认融合 默认 False,需显式开启 性能默认值与 PyTorch 不一致
load_state_dict 标准 + DCP 互操作 自定义 copy/distribute 路径 DCP 格式互操作待验证(见 distributed_checkpoint 模块)
get_model_state_dict DCP 标准生态 api.get_model_state_dict(未导出) 功能有,顶层可见性与 DCP 对齐待完善
disable_fsdp_gradient_division(Titan) set_gradient_divide_factor(1.0) set_gradient_scaling_factor + set_reduce_op_type 语义需按全局 token 数自行配置,非一键等价
fully_shard + 已 DTensor 参数 + 多维 mesh dp_mesh_dims 自动抽取 兼容模式推断 mesh 或要求 flatten CP+FSDP 迁移成本更高
MixedPrecisionPolicy 四字段 多 apply_grad_on_fp32_main_grad 🔶 FP32 main grad 为 Hyper 扩展,PyTorch 无直接字段

3.4 Hyper 有 · PyTorch 用户 API 无直接对标

接口 说明
HSDPModule 命名 + 动态混入 HSDP{OrigClass} 非 MRO FSDPModule;isinstance 检查需用 HSDPModule
replicate_params 声明式「不切分但 AR 同步」参数集合
comm_fusion + comm_fusion_zero_copy 显式融合通信与零拷贝 flat buffer(PT/MS 默认分化)
set_gradient_scaling_factor reduce 后乘法缩放,可 None 跳过热路径
set_reduce_op_type("avg"|"sum") 梯度归约类型(MS 路径 avg = SUM + div)
hsdp_sync_stream 等待异步梯度通信完成
load_state_dict 绕过 dispatch 支持 local-shard / global plain tensor 自动 distribute_tensor
zero_grad MS 路径 MindSpore scheduler 专用清零
MixedPrecisionPolicy.apply_grad_on_fp32_main_grad Torch 路径 FP32 main grad
MindSpore 双栈 enable_mindspore_backward_compat() + 平台分化 post_backward

四、顶层导出清单(hyper_parallel.__all__,本模块相关)

"fully_shard", "hsdp_sync_stream", "HSDPModule"

子模块 hyper_parallel.core.fully_shard:

"fully_shard", "HSDPModule"   # hsdp_sync_stream 仅在顶层 __init__

未导出但常用:

  • get_model_state_dict(core/fully_shard/api.py)
  • MixedPrecisionPolicy / OffloadPolicy / CPUOffloadPolicy(从 utils 经 fully_shard 参数类型引用)
  • FSDPMeshInfo / HSDPMeshInfo(内部调度)

策略类导入示例:

from hyper_parallel.core.fully_shard.utils import MixedPrecisionPolicy, CPUOffloadPolicy

五、维护说明

  • 本表为 §2 HSDP / FSDP 的接口索引;差距分析与 P0/P1 补齐路线见 #5。
  • 分布式 Checkpoint(hyper_parallel.core.distributed_checkpoint)与 FSDP 强相关但属 #1 独立子模块,本表仅列 get_model_state_dict 桥接项;DCP 全量表另开 Issue。
  • 源码变更导致接口漂移时,请同步更新本 Issue 总表。

父文档:#1 · 姊妹专项:#5 HSDP/FSDP 差距与路线图 · 关联:#7 DTensor 全量总表

likedislike