📎 上级文档: HyperParallel / PyTorch / TorchTitan 用户接口对标 #1 📎 姊妹专项(差距与路线图): HSDP / FSDP 对标 PyTorch #5
📎 上级文档: HyperParallel / PyTorch / TorchTitan 用户接口对标 #1
📎 姊妹专项(差距与路线图): HSDP / FSDP 对标 PyTorch #5
本 Issue 定位: 在 #1 模块划分下,汇总 §2 HSDP / FSDP 子模块的全量对外接口对标总表(表头:接口名 · 参数 · 功能 · hyper · pytorch · titan),作为 #5 差距梳理的接口索引附录。
分析依据(源码,2026-07):
hyper_parallel/__init__.py
hyper_parallel/core/fully_shard/
platform/*/fully_shard/
torch/distributed/fsdp/_fully_shard/
torch/distributed/_composable/replicate_with_fsdp.py
torchtitan/distributed/fsdp.py
fully_shard
图例: ✅ 有对标 · ⚠️ 部分对标 · ❌ 无 · 🔶 Hyper 扩展
列说明:
torch.distributed.fsdp
架构备注(影响对标方式): Hyper 通过动态混入 HSDP{OrigClass} 扩展模块,非 PyTorch @contract + MRO FSDPModule;类名对标为 HSDPModule ↔ FSDPModule。
HSDP{OrigClass}
@contract
FSDPModule
HSDPModule
module
list[module]
hyper_parallel.fully_shard
torch.distributed.fsdp.fully_shard
apply_fsdp_to_decoder
mesh
DeviceMesh | None
dp_mesh
edp_mesh
reshard_after_forward
bool
bool | int | None
get_fsdp_reshard_after_forward_policy
shard_placement_fn
Callable[[Parameter], Shard | None]
Shard | None
Shard | ShardPlacementResult | None
mp_policy
MixedPrecisionPolicy
apply_grad_on_fp32_main_grad
offload_policy
OffloadPolicy
CPUOffloadPolicy
cpu_offload=True
CPUOffloadPolicy()
ignored_params
set[Parameter]
replicate_params
replicate()
ReplicateModule
comm_fusion
False
FSDPParamGroup
comm_fusion_zero_copy
Optional[bool]
dp_mesh_dims
DataParallelMeshDims(shard=..., replicate=...)
isinstance(m, HSDPModule)
register_fsdp_forward_method
method_name
forward
torch.distributed.fsdp.register_fsdp_forward_method
share_comm_ctx
list[FSDPModule]
torch.distributed.fsdp.share_comm_ctx
get_cls_to_fsdp_cls
disable_fsdp_module_new_init
FSDPModule.__new__
replicate
torch.distributed._composable.replicate
disable_fsdp_gradient_division
unshard
async_op: bool = False
HSDPModule.unshard
FSDPModule.unshard
reshard
_UnshardHandle
UnshardHandle
_UnshardHandle.wait()
UnshardHandle.wait()
set_requires_gradient_sync
requires_gradient_sync
recurse=True
no_sync
set_requires_all_reduce
requires_all_reduce
set_reshard_after_forward
recurse=False
set_reshard_after_backward
reshard_after_backward
set_is_last_backward
is_last_backward: bool
set_modules_to_forward_prefetch
modules: list
set_modules_to_backward_prefetch
reset_iter_state
set_post_optim_event
event: torch.Event
set_gradient_divide_factor
factor: float
1.0
set_gradient_scaling_factor
factor: None | float | Tensor
set_reduce_op_type
"avg" | "sum"
gradient_divide_factor
set_custom_all_gather
comm: AllGather
set_custom_reduce_scatter
comm: ReduceScatter
set_all_reduce_hook
hook
stream=None
set_force_sum_reduction_for_comms
enable: bool
enable_fsdp_symm_mem
set_symm_mem_for_comm
backend="NCCL"
set_allocate_memory_from_process_group_for_comm
set_reduce_scatter_unused_params
reduce_scatter_unused_params
recurse
set_reduce_scatter_max_input_buffers
max_input_buffers
set_separate_reduce_scatter_group
enable
set_unshard_in_backward
unshard_in_backward: bool
load_state_dict
state_dict
strict=True
assign=False
nn.Module
zero_grad
param_dtype
reduce_dtype
output_dtype
cast_forward_inputs
pin_memory=True
DataParallelMeshDims
shard
edp_mesh_dims
FSDPMeshInfo
HSDPMeshInfo
DDPMeshInfo
shard_mesh_dim
replicate_mesh_dim
core/fully_shard/utils.py
_fsdp_common.py
_get_mesh_info
ShardPlacementResult
placement
mesh_info
_fsdp_common.ShardPlacementResult
AllGather
ReduceScatter
Comm
_fsdp_api.py
hsdp_sync_stream
hyper_parallel.hsdp_sync_stream
get_model_state_dict
model
options: StateDictOptions
core/fully_shard/api.py
__all__
torch.distributed.checkpoint.state_dict.get_model_state_dict
StateDictOptions
full_state_dict
cpu_offload
ignore_frozen_params
broadcast_from_rank0
get_hsdp_state
hsdp_utils.get_hsdp_state
fully_shard.state(module)
is_dtensor_managed_param
param
hsdp_utils
infer_fully_shard_param_mode
apply_gradient_scaling_factor
pp_enabled
reshard_after_forward_policy
ep_degree
enable_symm_mem
examples/torch/llama3/
policy: str
"always"
"never"
"default"
set_gradient_divide_factor(1.0)
本节从 总表 中筛出 hyper 列为 ❌ 或 ⚠️ 的项,按来源拆为四块。详细 P0/P1 补齐路线见 #5。
mesh.flatten()
reshard_after_forward: int | None
_fsdp_common
_composable.replicate
_fully_shard
Titan 不重新定义 FSDP 原语,底层均为 PyTorch fully_shard + FSDPModule。下表为 Titan 训练编排层 特有、Hyper 无对等公开 API 的项。
default
gradient_divide_factor=1.0
set_reshard_after_{forward,backward}
NotImplementedError
distributed_checkpoint
api.get_model_state_dict
isinstance
None
set_reduce_op_type("avg"|"sum")
avg
distribute_tensor
MixedPrecisionPolicy.apply_grad_on_fp32_main_grad
enable_mindspore_backward_compat()
post_backward
hyper_parallel.__all__
"fully_shard", "hsdp_sync_stream", "HSDPModule"
子模块 hyper_parallel.core.fully_shard:
hyper_parallel.core.fully_shard
"fully_shard", "HSDPModule" # hsdp_sync_stream 仅在顶层 __init__
未导出但常用:
utils
策略类导入示例:
from hyper_parallel.core.fully_shard.utils import MixedPrecisionPolicy, CPUOffloadPolicy
hyper_parallel.core.distributed_checkpoint
父文档:#1 · 姊妹专项:#5 HSDP/FSDP 差距与路线图 · 关联:#7 DTensor 全量总表
一、文档说明
本 Issue 定位: 在 #1 模块划分下,汇总 §2 HSDP / FSDP 子模块的全量对外接口对标总表(表头:接口名 · 参数 · 功能 · hyper · pytorch · titan),作为 #5 差距梳理的接口索引附录。
分析依据(源码,2026-07):
hyper_parallel/__init__.py、hyper_parallel/core/fully_shard/、platform/*/fully_shard/torch/distributed/fsdp/_fully_shard/、torch/distributed/_composable/replicate_with_fsdp.pytorchtitan/distributed/fsdp.py(编排层,底层调用 PyTorchfully_shard)图例: ✅ 有对标 · ⚠️ 部分对标 · ❌ 无 · 🔶 Hyper 扩展
列说明:
torch.distributed.fsdp)架构备注(影响对标方式): Hyper 通过动态混入
HSDP{OrigClass}扩展模块,非 PyTorch@contract+ MROFSDPModule;类名对标为HSDPModule↔FSDPModule。二、全量接口对标总表
fully_shardmodule/list[module]hyper_parallel.fully_shardtorch.distributed.fsdp.fully_shardapply_fsdp_to_decoder逐层调用fully_shard·meshDeviceMesh | Nonedp_mesh/edp_meshfully_shard·reshard_after_forwardbool(Hyper)/bool | int | None(PT)get_fsdp_reshard_after_forward_policy解析为 boolfully_shard·shard_placement_fnCallable[[Parameter], Shard | None](Hyper)Shard | NoneShard | ShardPlacementResult | Nonefully_shard·mp_policyMixedPrecisionPolicyapply_grad_on_fp32_main_grad🔶MixedPrecisionPolicy传入fully_shard·offload_policyOffloadPolicy/CPUOffloadPolicycpu_offload=True→CPUOffloadPolicy()fully_shard·ignored_paramsset[Parameter]fully_shard·replicate_paramsset[Parameter]replicate()/ReplicateModule)fully_shard·comm_fusionbool,默认FalseFSDPParamGroupfully_shard·comm_fusion_zero_copyOptional[bool]fully_shard·dp_mesh_dimsDataParallelMeshDims(shard=..., replicate=...)fully_shard· 返回值isinstance(m, HSDPModule))HSDP{OrigClass}FSDPModuleMROFSDPModuleregister_fsdp_forward_methodmodule,method_nameforward方法注册 pre/post FSDP hooktorch.distributed.fsdp.register_fsdp_forward_methodshare_comm_ctxlist[FSDPModule]torch.distributed.fsdp.share_comm_ctxget_cls_to_fsdp_clsdisable_fsdp_module_new_initFSDPModule.__new__特殊构造replicate/ReplicateModulemodule,mesh, …torch.distributed._composable.replicatedisable_fsdp_gradient_division注释兼容HSDPModule/FSDPModuleunshardasync_op: bool = FalseHSDPModule.unshardFSDPModule.unshardreshard_UnshardHandle/UnshardHandle_UnshardHandle.wait()UnshardHandle.wait()set_requires_gradient_syncrequires_gradient_sync,recurse=Trueno_sync(控制 RS+AR)set_requires_all_reducerequires_all_reduce,recurse=Trueset_reshard_after_forwardreshard_after_forward,recurse=Truerecurse=False未实现)set_reshard_after_backwardreshard_after_backward,recurse=Truerecurse=False未实现)set_is_last_backwardis_last_backward: boolset_modules_to_forward_prefetchmodules: listset_modules_to_backward_prefetchmodules: listreset_iter_stateset_post_optim_eventevent: torch.Eventset_gradient_divide_factorfactor: floatdisable_fsdp_gradient_division→1.0set_gradient_scaling_factorfactor: None | float | Tensorset_reduce_op_type"avg" | "sum"gradient_divide_factorset_custom_all_gathercomm: AllGatherset_custom_reduce_scattercomm: ReduceScatterset_all_reduce_hookhook,stream=Noneset_force_sum_reduction_for_commsenable: boolenable_fsdp_symm_mem内调用set_symm_mem_for_commbackend="NCCL"enable_fsdp_symm_memset_allocate_memory_from_process_group_for_commenable: boolset_reduce_scatter_unused_paramsreduce_scatter_unused_params,recurseset_reduce_scatter_max_input_buffersmax_input_buffers,recurseset_separate_reduce_scatter_groupenable,recurseset_unshard_in_backwardunshard_in_backward: boolload_state_dictstate_dict,strict=True,assign=Falsenn.Module+ DCPzero_gradnn.ModuleMixedPrecisionPolicyparam_dtype,reduce_dtype,output_dtype,cast_forward_inputsapply_grad_on_fp32_main_gradOffloadPolicyCPUOffloadPolicypin_memory=Truecpu_offload=TrueDataParallelMeshDimsshard,replicate(mesh 轴名)dp_mesh_dims/edp_mesh_dimsFSDPMeshInfo/HSDPMeshInfo/DDPMeshInfomesh,shard_mesh_dim,replicate_mesh_dimcore/fully_shard/utils.py(内部)_fsdp_common.py(内部)_get_mesh_infoShardPlacementResultplacement,mesh_info_fsdp_common.ShardPlacementResultAllGather/ReduceScatter/Comm_fsdp_api.pyhsdp_sync_streamhyper_parallel.hsdp_sync_streamget_model_state_dictmodel,options: StateDictOptionscore/fully_shard/api.py(未进__all__)torch.distributed.checkpoint.state_dict.get_model_state_dictStateDictOptionsfull_state_dict,cpu_offload,ignore_frozen_params,broadcast_from_rank0StateDictOptionsget_hsdp_statemodulehsdp_utils.get_hsdp_statefully_shard.state(module)is_dtensor_managed_paramparamhsdp_utilsinfer_fully_shard_param_modehsdp_utilsapply_gradient_scaling_factorhsdp_utilsapply_fsdp_to_decodermodel,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_memexamples/torch/llama3/参考)torchtitan/distributed/fsdp.pyget_fsdp_reshard_after_forward_policypolicy: str,pp_enabled"always"/"never"/"default"→ booldisable_fsdp_gradient_divisionmodelset_gradient_divide_factor(1.0)set_gradient_scaling_factorFSDPModuleenable_fsdp_symm_memmodelFSDPModule三、差距速查(与 #5 联动)
3.1 仅 PyTorch 有 · Hyper 无
fully_shard·dp_mesh_dimsDataParallelMeshDimsmesh.flatten()reshard_after_forward: int | Nonefully_shardreshard_after_forward_policy部分语义shard_placement_fn→ShardPlacementResult_fsdp_commonregister_fsdp_forward_methodtorch.distributed.fsdpforward方法 FSDP hookshare_comm_ctxreset_iter_stateFSDPModuleset_post_optim_eventFSDPModuleset_custom_all_gather/set_custom_reduce_scatterFSDPModule+CommABCset_all_reduce_hookFSDPModuleset_symm_mem_for_comm/set_force_sum_reduction_for_commsFSDPModuleset_allocate_memory_from_process_group_for_commFSDPModuleset_reduce_scatter_unused_paramsFSDPModuleset_reduce_scatter_max_input_buffersFSDPModuleset_separate_reduce_scatter_groupFSDPModuleset_unshard_in_backwardFSDPModulereplicate/ReplicateModule_composable.replicatereplicate_params集合get_cls_to_fsdp_cls/disable_fsdp_module_new_init_fully_shard3.2 仅 TorchTitan 有 · Hyper 无
apply_fsdp_to_decodertorchtitan/distributed/fsdp.pyget_fsdp_reshard_after_forward_policydefault= PP 时默认不 resharddisable_fsdp_gradient_divisiongradient_divide_factor=1.0enable_fsdp_symm_memapply_fsdp_to_decoder内 EP prefetch 链3.3 PyTorch 有 · Hyper 有但语义偏弱(⚠️)
reshard_after_forwardbool | int | Noneboolset_reshard_after_forward变通shard_placement_fnShardPlacementResult(含 mesh_info)Shard | Noneset_reshard_after_{forward,backward}·recurse=FalseNotImplementedErrorrecurse=True等价实现set_requires_gradient_sync·recursecomm_fusionFalse,需显式开启load_state_dictdistributed_checkpoint模块)get_model_state_dictapi.get_model_state_dict(未导出)disable_fsdp_gradient_division(Titan)set_gradient_divide_factor(1.0)set_gradient_scaling_factor+set_reduce_op_typefully_shard+ 已 DTensor 参数 + 多维 meshdp_mesh_dims自动抽取MixedPrecisionPolicyapply_grad_on_fp32_main_grad🔶3.4 Hyper 有 · PyTorch 用户 API 无直接对标
HSDPModule命名 + 动态混入HSDP{OrigClass}FSDPModule;isinstance检查需用HSDPModulereplicate_paramscomm_fusion+comm_fusion_zero_copyset_gradient_scaling_factorNone跳过热路径set_reduce_op_type("avg"|"sum")avg= SUM + div)hsdp_sync_streamload_state_dict绕过 dispatchdistribute_tensorzero_gradMS 路径MixedPrecisionPolicy.apply_grad_on_fp32_main_gradenable_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五、维护说明
hyper_parallel.core.distributed_checkpoint)与 FSDP 强相关但属 #1 独立子模块,本表仅列get_model_state_dict桥接项;DCP 全量表另开 Issue。父文档:#1 · 姊妹专项:#5 HSDP/FSDP 差距与路线图 · 关联:#7 DTensor 全量总表