FSDP 梯度通信处理理解
这里按当前 fully_shard 对 DTensor 的两条路径梳理一下梯度通信 group 的推导规则。
DTensor Compact / Compat
DTensor 参数
_spmd_mesh = 原 DTensor mesh_spmd_placements = 原 DTensor placements- 不做 FSDP
ReduceScatter AllReducegroup 来自原 placements 中所有Replicate()轴Shard轴即使对应 mesh dim size 是 1,也不会进入AllReducegroup
普通 Tensor / LocalParam
_spmd_mesh = compat_mesh,也就是 scheduler 从 DTensor 参数里重新扫描出来的 mesh_spmd_placements = 全 Replicate- 不做
ReduceScatter AllReducegroup 是compat_mesh上所有 replicate axes 组合出来的 group。多维 mesh 时不是只用 dim0,而是可能覆盖整个compat_mesh- 这里隐含要求:开发者要保证
LocalParam在这个compat_mesh的 replicated group 上初始化一致。否则梯度虽然同步了,但参数初始值本身已经不一致
DTensor Unified
DTensor 参数
_spmd_mesh = DeviceMesh.concatenate([FSDP/HSDP mesh, 原 DTensor mesh])- placements 先是
[Replicate for FSDP/HSDP mesh dims] + 原 DTensor placements - 然后 FSDP shard dim 被替换成
Shard(...)或StridedShard(...) ReduceScatter在显式传入的 FSDP/HSDP shard 轴上做AllReducegroup 来自最终 unified placements 里的Replicate()轴,并排除 FSDP shard 轴- 因此它会包含 HSDP replicate 轴,也会包含原 DTensor 自带的
Replicate()轴 - 但不会包含原 DTensor 的
Shard轴,包括 size=1 的 shard 轴
普通 Tensor / LocalParam
_spmd_mesh = 显式传入的 FSDP/HSDP mesh- 初始 placements 是全
Replicate - 如果启用 FSDP shard,则 shard dim 被替换为
Shard - 1D FSDP:通常只有
ReduceScatter,没有AllReduce - 2D HSDP:shard dim 做
ReduceScatter,replicate dim 做AllReduce - 不会拼接 DTensor mesh,也不会从 DTensor placements 推导任何东西
- 同样要求
LocalParam在被认为 replicated 的维度上初始化一致;并且 FSDP shard 前每个 rank 的 full param 也最好一致,否则 shard 后拼出来的是不同 rank 初始参数的混合体
关键代码链路
infer_fully_shard_param_mode():决定LOCAL_PARAM / DTENSOR_COMPAT / DTENSOR_UNIFIEDTorchHSDPScheduler._new_cell_state():mesh=None时扫描 DTensor mesh,构造DDPMeshInfoHSDPParam._get_base_spmd_placements():决定_spmd_mesh和初始 placementsHSDPParam._apply_data_parallel_placements():插入 FSDP/HSDP shard placementHSDPParam._init_group_infos():分别构造ReduceScattergroup 和AllReducegroupHSDPParam._build_layout_driven_group_info():只从placement.is_replicate()轴推导AllReducegroup


replicate_params 的处理路径补充
这里单独梳理一下 replicate_params。它的核心语义不是“走普通 FSDP shard 但保留副本”,而是:
- 初始化时
enable_fsdp_shard=False - 参数仍然被包装成
HSDPParam - 但不会插入 FSDP shard placement
- 正向仍走
unshard / wait_for_unshard / shard状态机 - 实际不会发起 FSDP
AllGather,因为is_sharded=False,_get_unsharded_param_data()里直接本地 copy - 反向不做真实
ReduceScatter;梯度规约主要依赖 layout-drivenAllReduce
初始化与正向
关键链路:
TorchHSDPStateV2._init_hsdp_params()/MindSporeHSDPStateV2._init_hsdp_params()replicate_params = set(config.replicate_params or ())enable_fsdp_shard = param not in replicate_params- 命中
replicate_params的参数会进入self.replicate_params
HSDPParam._init_sharded_param()- 因为
uses_param_shard == False,所以shard_world_size=1 is_sharded=Falsesharded_group_info变成 invalid / rank_size=1unsharded_group_info仍通过_build_layout_driven_group_info()从最终 placements 推导
- 因为
HSDPState.unshard(unshard_replicate=True)- replicate params 也会调用
param.unshard(async_op)
- replicate params 也会调用
HSDPParam._get_unsharded_param_data()not self.is_sharded时只分配 world_size=1 的 output buffer,然后 copy 本地 tensor- 不调用
dist.all_gather_into_tensor
所以严格说,replicate_params 正向不是“做 AllGather”,而是“经过 AllGather/unshard 这套代码路径,但由于未分片,通信退化成本地 copy”。
网络中没有 DTensor
这时 param_mode=LOCAL_PARAM。
placements 和 group
_spmd_mesh = mesh_info.mesh_spmd_placements = 全 Replicate- 因为
enable_fsdp_shard=False,_apply_data_parallel_placements()不会把 shard dim 替换成Shard(...) _build_layout_driven_group_info()会把所有Replicate()轴都放进AllReducegroup
结果:
- 默认无 DTensor 且用户不传 mesh 时,
fully_shard()会创建 1D world mesh,replicate_params的梯度会在整个 1D mesh 上AllReduce - 用户显式传 1D FSDP mesh 时,同样是在这个 1D mesh 上
AllReduce - 用户显式传 2D HSDP mesh 时,由于 replicate param 没有 shard dim,两个 mesh 轴都还是
Replicate(),因此AllReducegroup 会覆盖所有 replicate axes,语义上接近在整个 2D mesh 上规约,而不是只在 HSDP replicate dim 上规约
反向通信路径
Torch non-comm-fusion 路径:
- 入口:
TorchHSDPStateV2.post_backward() replicate_params会被_iter_managed_params()遍历到- 在
_issue_reduce_scatter_for_current_module()中进入 params_to_reduce reduce_scatter_grad()因为is_sharded=False不会发起真实ReduceScatter,只做本地 pack/copy,并把结果作为 reduce output- 如果
_should_run_all_reduce(hsdp_param)为真,即dp_size > 1,会按unsharded_group_info.group构造AllReduceParamGroup - 后续
_wait_prev_reduce_scatter()等待这个 fake RS,_issue_prev_fused_allreduce()发起异步AllReduce - 最终在 root backward 的
delay_apply_reduce_grads()里 wait 并 apply 到 sharded/local param grad
Torch comm-fusion 路径:
replicate_params不进入HSDPParamGroup,因为_comm_fusion_unsupported_reason()对enable_fsdp_shard=False返回 unsupportedpost_backward_for_comm_fusion()先处理 sharded params 的 fused pipeline- 然后单独遍历
self.replicate_params - 调用
_queue_compat_all_reduce()发起纯AllReduce - wait/apply 由下一轮或 root hook 前的
reduce_params()处理
MindSpore 路径:
- MindSpore 显式抽了
_queue_replicate_params_allreduce() dp_size > 1时走_queue_compat_all_reduce()dp_size <= 1时走_apply_pending_unsharded_grad_locally()- comm-fusion 和 non-comm-fusion 都会在合适位置调用
_queue_replicate_params_allreduce()
这里 Torch 与 MindSpore 有一点实现差异:MindSpore 对 replicate_params 的 no-comm local apply 路径更显式;Torch comm-fusion 路径在 _queue_compat_all_reduce() 中如果 dp_size <= 1 会直接 return,这块需要结合实际配置确认是否有遗漏。
网络中有 DTensor
这里要区分 DTENSOR_COMPAT 和 DTENSOR_UNIFIED,同时还要区分 replicate param 自己是不是 DTensor。
DTensor Compact / Compat
触发条件:网络中存在 DTensor,用户没有显式传 FSDP/HSDP mesh。
Scheduler 会从 DTensor 参数扫描出 compat_mesh,构造:
mesh_info = DDPMeshInfo(mesh=compat_mesh, replicate_mesh_dim=0)
但每个参数自己的 param_mode 仍按该参数是否是 DTensor 推导。
replicate DTensor 参数:
_spmd_mesh = 原 DTensor mesh_spmd_placements = 原 DTensor placements- 不插入 FSDP shard placement
- 正向无真实
AllGather - 反向无真实
ReduceScatter AllReducegroup 只来自原 placements 中的Replicate()轴- 原 DTensor 的
Shard轴不会参与AllReduce,即使对应 mesh dim size 是 1
replicate LocalParam / 普通 Tensor:
_spmd_mesh = compat_mesh_spmd_placements = 全 Replicate- 正向无真实
AllGather - 反向无真实
ReduceScatter AllReducegroup 覆盖compat_mesh上所有 replicate axes- 这隐含要求 LocalParam 在这个 group 上初始化一致,否则梯度同步也无法修复初始参数不一致
DTensor Unified
触发条件:网络中存在 DTensor,且用户显式传入 FSDP/HSDP mesh。
replicate DTensor 参数:
_spmd_mesh = DeviceMesh.concatenate([FSDP/HSDP mesh, 原 DTensor mesh])- 初始 placements 是
[Replicate for FSDP/HSDP mesh dims] + 原 DTensor placements - 因为
enable_fsdp_shard=False,FSDP shard dim 不会被替换成Shard(...) - 所以显式传入的 FSDP/HSDP mesh 轴都会保留为
Replicate() AllReducegroup 来自最终 placements 中所有Replicate()轴- 因此 group 会包含显式 FSDP/HSDP mesh 的所有轴,也会包含原 DTensor 自带的
Replicate()轴 - 原 DTensor 的
Shard轴仍不会参与AllReduce
replicate LocalParam / 普通 Tensor:
_spmd_mesh = 显式传入的 FSDP/HSDP mesh_spmd_placements = 全 Replicate- 不拼接 DTensor mesh
- 不从其它 DTensor 的 placements 推导任何东西
- 1D FSDP mesh:在这个 1D mesh 上
AllReduce - 2D HSDP mesh:因为两个轴都保留
Replicate(),AllReducegroup 会覆盖所有 replicate axes,通常就是整个 2D mesh
reduce op、混合精度和 offload
_resolve_default_reduce_op():如果任意 managed param 是DTENSOR_COMPAT / DTENSOR_UNIFIED,默认 reduce op 使用SUM;否则默认AVGall_reduce_grad(dtype=self._reduce_dtype, ...)会按 mixed precision 的reduce_dtype转换通信 bufferapply_reduced_grad()再根据mp_policy.apply_grad_on_fp32_main_grad写入grad或main_grad- CPU offload 在
apply_reduced_grad()中处理,必要时把 reduced grad 搬到 CPU,并在 state 层做 stream synchronize
关键函数索引
TorchHSDPStateV2._init_hsdp_params()/MindSporeHSDPStateV2._init_hsdp_params():识别replicate_params并设置enable_fsdp_shard=FalseHSDPParam._get_base_spmd_placements():决定 DTensor/LocalParam 的基础 mesh 和 placementsHSDPParam._apply_data_parallel_placements():普通 sharded param 会插入 shard placement;replicate_params不会HSDPParam._init_group_infos():构造sharded_group_info与unsharded_group_infoHSDPParam._build_layout_driven_group_info():只从placement.is_replicate()轴推导AllReducegroupHSDPParam._get_unsharded_param_data():replicate_params正向 unshard 退化成本地 copyHSDPParam.reduce_scatter_grad():is_sharded=False时反向 RS 退化成本地 pack/copyHSDPParam.all_reduce_grad():纯AllReduce或 RS 后接 AR 的实际通信入口TorchHSDPStateV2._issue_reduce_scatter_for_current_module():Torch non-comm-fusion 下 replicate params 被纳入这里TorchHSDPStateV2.post_backward_for_comm_fusion():Torch comm-fusion 下 replicate params 走单独纯 AR side pathMindSporeHSDPStateV2._queue_replicate_params_allreduce():MindSpore replicate params 的显式处理入口


replicate_params 与 shard_size=1 路径归一化审视
这次可以把 replicate_params 和 fully_shard 中 shard_size == 1 的参数统一看成 no-shard param:
- 没有真实 FSDP 参数切分
- 正向不需要真实
AllGather - 反向不需要真实
ReduceScatter - 只需要根据最终 layout 推导出的 replicate group 做
AllReduce - 如果 replicate group size 也是 1,则连
AllReduce都不需要,本地 grad 直接落到目标 grad
当前代码现状
replicate_params
初始化入口:
TorchHSDPStateV2._init_hsdp_params()MindSporeHSDPStateV2._init_hsdp_params()
命中 config.replicate_params 后:
enable_fsdp_shard=False- 参数进入
self.replicate_params HSDPParam.uses_param_shard == False_init_sharded_param()中shard_world_size=1is_sharded=Falsesharded_group_info是 invalid / rank_size=1unsharded_group_info仍由_build_layout_driven_group_info()从Replicate()placement 轴推导
shard_size == 1 的普通 fully_shard 参数
这类参数通常仍然是:
enable_fsdp_shard=True- 参数在
self.hsdp_params / self.sharded_hsdp_params - 但由于 mesh shard dim size 为 1,或实际
sharded_group_info.rank_size == 1 is_sharded=False
所以从通信语义上看,它和 replicate_params 很接近:不需要真实 AG/RS,只需要按 replicate group 做 AR 或本地 apply。
正向路径:当前多了一次 copy
当前 HSDPState.unshard() / prefetch() 对 replicate params 和普通 sharded params 都会调用 param.unshard(async_op)。
在 HSDPParam._get_unsharded_param_data() 中:
- 如果
not self.is_sharded - 会分配
world_size=1的all_gather_outputs - 然后把
all_gather_inputcopy 到all_gather_outputs[0] wait_for_unshard()再调用init_unsharded_param()init_unsharded_param()从all_gather_outputs[0]unpack 出_unsharded_param- 最后
to_unsharded()把_unsharded_param绑定回 module
也就是说,shard_size == 1 时确实没有真实 AllGather,但仍然有一次本地 buffer copy。这个 copy 对无切分边界场景来说不是必要的。
优化方向:不要简单等价成 unsharded_param = sharded_param
这个优化方向成立,但要注意:当前 sharded_param 通常是 Parameter(DTensor),而 forward 期 module 期望看到的对象不一定就是这个 sharded_param。
更准确的目标应该是:
no-shard param 在 unshard 阶段构造一个“计算视图”,这个计算视图尽量与 sharded storage 共享底层存储,避免 copy 和 all_gather buffer。
需要分情况:
LOCAL_PARAM
原始普通 Tensor 参数在 forward 期通常期望是本地 Tensor Parameter。
当前 sharded_param 是 DTensor.from_local(...) 包出来的参数。如果直接把 unsharded_param = sharded_param 绑定进 module,会让普通 Tensor 参数的 forward 看到 DTensor,语义会变。
所以 LOCAL_PARAM 更合理的 fast path 是:
_unsharded_param = Parameter(local_tensor_view)local_tensor_view与sharded_param.to_local()/_sharded_local_tensor共享存储- 不经过
all_gather_outputs - 不 copy
DTENSOR_COMPAT
如果原参数本来就是 DTensor,并且 compat 路径下:
_spmd_mesh = 原 DTensor mesh_spmd_placements = 原 DTensor placements- no-shard 时 local shape 不变
这种场景最接近可以直接复用 sharded_param。但仍要确认它暴露给 forward 的 DTensor layout 与原始 DTensor layout 完全一致。
DTENSOR_UNIFIED
Unified 下即使 shard_size == 1,sharded_param 的 layout 也可能是:
DeviceMesh.concatenate([FSDP/HSDP mesh, 原 DTensor mesh])- placements 带有 FSDP/HSDP mesh 前缀
而 forward 期更自然的语义是原 DTensor view。因此这里不能简单把 sharded_param 直接绑定给 module。更合适的是:
- 用同一个 local tensor 构造原 DTensor mesh/placements 的 compute view
- 也就是类似
DTensor.from_local(local_tensor, orig_dtensor_mesh, orig_dtensor_placements) - 但避免 all_gather buffer 和本地 copy
混合精度
如果 param_dtype is not None,那 forward compute 参数需要 cast 到目标 dtype。
这时不能与 sharded_param 完全同对象同 storage:
unsharded_param需要是 cast 后的 compute Parameter- 梯度会挂在
unsharded_param.grad - 后续 reduce/apply 仍然要把 grad 归约后写回
sharded_param.grad或main_grad
所以 no-copy alias fast path 只适合 param_dtype is None 的场景;混合精度下可以跳过 AG buffer,但 dtype cast 本身仍然会产生新 tensor。
prefetch / wait_for_unshard
建议给 HSDPParam 加一个清晰的判断,例如:
requires_all_gather = self.is_shardedrequires_reduce_scatter = self.is_shardedis_no_shard_param = not self.is_sharded
在 unshard(async_op=True) / prefetch() 中:
- no-shard param 不分配
all_gather_outputs - 不调用
all_gather_into_tensor prefetch_handle=None- 只准备 compute view
wait_for_unshard()直接把 compute view 绑定回 module
这样 prefetch 对 no-shard param 就退化为“恢复 module 参数绑定”,而不是“准备一个 world_size=1 的 all_gather buffer”。
reshard / free_unsharded_param
如果 no-shard fast path 下:
- 无混合精度
- 无 offload
- compute view 与 sharded storage 共享存储,甚至是同一个对象
那么 to_sharded() 不能无脑:
- copy
_unsharded_param -> sharded_param free_unsharded_param()- resize 掉仍在使用的 storage
建议给 param 增加一个标志,例如:
_unsharded_param_aliases_sharded_storage- 或
_unsharded_param_is_sharded_param
然后:
- 如果 alias 同 storage,
to_sharded()只需要把 module 绑定回sharded_param - 不需要 copy
free_unsharded_param()不应释放 alias 依赖的 storage_unsharded_param是否置空要谨慎:如果是同对象,可以保留;如果是单独 compute view,可以清理 wrapper 但不能释放底层 sharded storage
梯度处理
no-shard param 的理想规约路径应该是:
reduce_scatter_group = None / invalid- 不进入真实
ReduceScatter - 如果
unsharded_group_info.rank_size > 1,直接在这个 group 上AllReduce - 如果
rank_size == 1,本地 apply
这里要特别小心 同对象 alias 的梯度重复累加问题。
如果 unsharded_param is sharded_param,那么 forward 产生的 grad 可能已经在 sharded_param.grad 上。此时如果再调用当前的 apply_reduced_grad(),它会看到目标 sharded_param.grad 已经存在,然后执行累加逻辑,可能把同一个 grad 加两次。
因此 no-shard alias 路径需要单独的 grad-finalize 逻辑:
- 无 AR、无 dtype cast:
sharded_param.grad已经是最终 grad,只需标记/清理状态 - 需要 AR:可以直接对
sharded_param.grad的 local tensor 原地AllReduce,wait 后不要再通过apply_reduced_grad()追加一次 - 需要 reduce_dtype cast / FP32 main_grad:用临时 reduced buffer,wait 后再写入
main_grad或目标 grad - DTensor grad 仍要走
_to_local_unsharded_grad()/ partial reduce / redistribute 的规范化逻辑
当前 MindSpore 路径已经更接近这个方向:
_should_skip_reduce_scatter_issue()会跳过hsdp_param.shard_size <= 1post_backward()中shard_size <= 1时直接_queue_compat_all_reduce()或_apply_pending_unsharded_grad_locally()replicate_params有专门的_queue_replicate_params_allreduce()
Torch 路径当前更混杂:
- non-comm-fusion 下
_issue_reduce_scatter_for_current_module()会把 no-shard 参数也收进去 reduce_scatter_grad()内部因为is_sharded=False退化成本地 pack/copy- 如果还需要 AR,再通过
AllReduceParamGroup继续走 fake-RS + AR 流程
这块可以重构为与 MindSpore 更一致:
- 先统一识别 no-shard param
- no-shard param 不进入 RS pipeline /
HSDPParamGroup - 统一走
queue_no_shard_grad_reduce() - 内部根据
unsharded_group_info.rank_size决定 AR 或 local apply
offload
你提到“切分份数 == 1 的参数不用 offload,因为没有冗余”,这个方向和 no-copy fast path 是一致的:如果 no-shard param 仍按 CPU offload 处理,那么 forward 前仍然要 CPU -> device copy,fast path 的收益会被抵消。
不过这里是行为语义变更,需要明确策略:
- 如果 CPUOffloadPolicy 的语义是“只 offload FSDP shard 存储”,那么 no-shard param 可以排除 offload
- 如果用户把 CPUOffloadPolicy 理解成“所有 managed params 都可 offload 以省 device memory”,那跳过 no-shard offload 会改变显存行为
代码上如果选择 no-shard 不 offload,需要同步调整:
_validate_cpu_offload_params():不要要求 no-shard param 必须在 CPU_init_sharded_param():no-shard param 不转 CPU / pin memoryall_gather_inputs:no-shard fast path 不再通过 offload input 触发 device copyapply_reduced_grad():no-shard alias grad 不应再被搬回 CPU
我的建议是第一阶段先保守:
- no-shard fast path 在
offload_to_cpu=False时启用 offload_to_cpu=True时先保留当前路径
等语义确认后,再决定是否把 no-shard param 从 offload policy 中排除。
建议的重构落点
可以抽一个统一分支,而不是继续让 replicate_params 和 shard_size == 1 分散在不同路径里:
- 在
HSDPParam上增加属性:is_no_shard_param = not self.is_shardedrequires_all_gather = self.is_shardedrequires_reduce_scatter = self.is_sharded
- 抽
prepare_unsharded_compute_param():- sharded 参数:走现有 AG/unpack
- no-shard LOCAL_PARAM:构造 local tensor compute view,尽量共享 storage
- no-shard DTENSOR_COMPAT:可在 layout 完全一致时复用/alias
- no-shard DTENSOR_UNIFIED:构造原 DTensor layout 的 compute view,避免暴露 unified layout
- mixed precision:cast 后创建 compute Parameter
- 抽
finalize_no_shard_grad():- rank_size > 1:layout-driven
AllReduce - rank_size == 1:local apply 或直接保留 grad
- alias 同对象时避免
apply_reduced_grad()双重累加
- rank_size > 1:layout-driven
- Torch 对齐 MindSpore:
shard_size <= 1不进入 RS pipelinereplicate_params和 no-shard normal params 走同一条 pure AR/local apply 路径
- comm_fusion:
- no-shard param 不进入
HSDPParamGroup - 作为 side path 做 pure AR/local apply
- no-shard param 不进入
这个重构的核心不是“所有场景都让 unsharded_param = sharded_param”,而是“所有 no-shard 场景都不再构造 world_size=1 的 AG buffer,不再 fake ReduceScatter,并且在保持 forward 期参数语义不变的前提下尽量 alias storage”。


FSDP 代码重构 与 dim-0 非均匀切分 RFC
1. 基本信息与阅读约定
coreMengXiangyu/fully_shard、platform/torch/fully_shard、platform/mindspore/fully_shard、core/dtensor、core/distributed_checkpointgrad_comm_overlap.md业界实现分析
_fully_shard.py2. 背景、目标与总体决策
2.1 当前问题
FullyShardParamMode包含LOCAL_PARAM、DTENSOR_COMPAT、DTENSOR_UNIFIEDfully_shard()只计算any(is_dtensor_managed_param(...))MeshInfo所有权过粗HSDPState持有 unit 级MeshInfo,param/param-group 再从 state 侧信息推导通信replicate_params与普通 shard 参数难以拥有不同 DP 通信语义,TP group 也容易被混入 FSDP/HSDP group 决策hsdp_params、sharded_hsdp_params、replicate_params,以及is_shard、is_replicate_shardunshard_replicate等配套 flagshard_size == 1或replicate_params仍创建all_gather_output并 copyReduceScatterPlan、Hyper DTensor global shape、DCP offset 均存在整除假设准确类名是
DTENSOR_COMPAT,本文不使用DTensor_Compact等错误写法。2.2 目标
DTENSOR_COMPAT、DTENSOR_UNIFIED对生命周期和梯度通信的控制。MeshInfo从HSDPState下沉到每个HSDPParam,且只表达 FSDP/HSDP 的 shard/replicate 通信;HSDPParam从原 Parameter DTensor placements 提供 TP 通信元数据,TP AR 只由根反向钩子在 FSDP/HSDP/DP 完成后触发,不再将 HSDP AR 与 TP AR 合并成笛卡尔积 group 做一次通信。replicate_params与shard_size == 1归一为无需 AllGather 的参数生命周期。2.3 非目标
3. Hook 触发顺序与状态机
本章先回顾“FSDP 流程如何被触发”。参数如何 unshard、梯度如何通信分别在第 5、6 章展开。
3.1 Hook 注册与宏观调用链
[Hyper 当前] Torch 与 MindSpore 都由 scheduler 注册 forward pre/post hook;forward 输入经
PostBackwardFunction包装,forward 输出注册 backward pre-hook。输出 grad hook 还会把 root final callback 放入 autograd 引擎队列。flowchart LR FP[forward pre-hook] --> U[unshard] U --> F[module forward] F --> FO[forward post-hook] FO --> R1[可选 reshard] R1 --> OB[output grad hook] OB --> UB[backward unshard] OB --> QC[queue root final callback] UB --> BW[module backward] BW --> PB[PostBackwardFunction.backward<br/>仅当输入 requires_grad] PB --> POST[post_backward] QC --> ROOT[root final callback<br/>补 post_backward + drain]PostBackwardFunction是否执行取决于该 FSDP unit 的输入中是否存在requires_grad=TrueTensor:PostBackwardFunction,root callback 负责补执行post_backward()。PostBackwardFunction.backward()可能先把 state 置为BACKWARD,但这只表示 post-backward hook 已触发,不表示所有异步通信已完成。scheduler_state == BACKWARD推断“无需收尾”。3.2 状态定义
NonePRE_FORWARDFORWARDPRE_BACKWARDBACKWARDPostBackwardFunction或 root fallbackpost_backward()已触发;pending reduction 仍可能未完成3.3 正常 forward → backward
sequenceDiagram participant A as User/Autograd participant S as HSDPScheduler participant P as HSDPState A->>S: forward_pre [None -> PRE_FORWARD] S->>P: unshard() A->>S: forward_post [PRE_FORWARD -> FORWARD] S->>P: reshard_after_forward A->>S: output grad hook [FORWARD -> PRE_BACKWARD] S->>P: unshard(backward) S->>A: queue root final callback A->>S: PostBackwardFunction.backward(若输入可求导) S->>P: post_backward [PRE_BACKWARD -> BACKWARD] A->>S: root final callback S->>P: 补遗漏的 post_backward S->>P: 无条件 drain pending reductions实测关键顺序:
3.4 不可重入重计算
不可重入 checkpoint 的 backward 先触发输出 grad hook,状态已经进入
PRE_BACKWARD,随后 checkpoint 重算再次进入 forward pre-hook。该 hook 必须幂等跳过;early-stop 下重算 forward post-hook 可能根本不触发。sequenceDiagram participant C as Non-reentrant checkpoint participant S as HSDPScheduler participant P as HSDPState C->>S: 原 forward_pre [None -> PRE_FORWARD] S->>P: unshard C->>S: 原 forward_post [PRE_FORWARD -> FORWARD] S->>P: reshard C->>S: output grad hook [FORWARD -> PRE_BACKWARD] S->>P: backward unshard C->>S: 重算 forward_pre [PRE_BACKWARD -> PRE_BACKWARD] Note over C,S: 幂等跳过;forward_post 可能因 early-stop 不出现 C->>S: root final callback S->>P: post_backward [PRE_BACKWARD -> BACKWARD] S->>P: final drain3.5 可重入重计算
可重入 checkpoint 的第一次 forward 在
no_grad下运行;backward 时完整重算 forward,所以 backward 期间会再次出现一对 forward hook。sequenceDiagram participant C as Reentrant checkpoint participant S as HSDPScheduler participant P as HSDPState C->>S: no_grad 原 forward_pre [None -> PRE_FORWARD] S->>P: unshard C->>S: no_grad 原 forward_post [PRE_FORWARD -> FORWARD] S->>P: reshard C->>S: backward 重算 forward_pre [FORWARD -> PRE_FORWARD] S->>P: unshard C->>S: 重算 forward_post [PRE_FORWARD -> FORWARD] S->>P: reshard C->>S: output grad hook [FORWARD -> PRE_BACKWARD] S->>P: backward unshard C->>S: root final callback S->>P: post_backward [PRE_BACKWARD -> BACKWARD] S->>P: final drain3.6 测试与日志证据
现有真实两 rank CPU/Gloo autograd 探针:
tests/torch/fully_shard/_test_fully_shard_hook_state_machine.pytests/torch/fully_shard/test_fully_shard_hook_state_machine.py已执行命令:
结果:两个 rank 均
1 passed。日志:4. 初始化层拦截:只允许全 DTensor 或全普通 Tensor
4.1 唯一的参数类别拦截
[Hyper 当前]
fully_shard()在过滤 ignored/already-managed 参数后只计算has_dtensor_param = any(...),没有拒绝混用。[目标] 在同一位置增加一次、且仅一次 managed-param 同构校验:
replicate_params仍属于 managed params,不能绕过该检查。嵌套 FSDP unit 各自校验,父 unit 不重复检查子 unit 已管理的参数。后续代码不得再次按单参数推导LOCAL_PARAM/DTENSOR_COMPAT/DTENSOR_UNIFIED。4.2 全 DTensor 的初始化约束
[PyTorch 2.9]
fully_shard(mesh=None)创建 default process group 上的 1-D WORLD mesh,不会从 DTensor 参数的 TP mesh 推导 DP mesh。FSDP mesh 与 TP/EP mesh 必须:标准二维 TP 场景中,
mesh=None创建的 WORLD mesh 与 TP root mesh 不同,PyTorch 在初始化阶段失败,不会进入 forward/backward,也不能声称会在 TP 轴执行重复通信。[目标]:
mesh=None保留默认 WORLD FSDP mesh 行为。mesh=None在 API 边界ValueError。Partialplacement 时本期拒绝。正确调用:
root_mesh = init_device_mesh( "npu", (dp_size, tp_size), mesh_dim_names=("dp", "tp"), ) parallelize_module(module, root_mesh["tp"], parallelize_plan) fully_shard(module, mesh=root_mesh["dp"])HSDP + TP 使用
("replicate", "fsdp", "tp")三维 root mesh,传给fully_shard()的是root_mesh[("replicate", "fsdp")]。错误调用:
parallelize_module(module, root_mesh["tp"], parallelize_plan) fully_shard(module, mesh=None) # DTensor 参数不能隐式推导 DP mesh fully_shard(module, mesh=root_mesh["tp"]) # 不能复用 TP 轴作为 FSDP 轴初始化拦截只负责输入合法性和静态元数据构造,不参与每次 unshard、reshard 或 backward 的动态分支。
5. Unshard/reshard 流程重构
5.1 [Hyper 当前] 参数集合、状态与 flag
hsdp_paramssharded_hsdp_paramsreplicate_paramsis_shardsharded_hsdp_params当前是否处于 sharded 状态is_replicate_shardreplicate_params的第二套名义 sharded 状态unshard_replicateunshard()是否处理 replicate 参数的对象/storage 切换shard_replicateshard()是否把 replicate 参数切回持久对象wait_for_replicatewait_for_unshard()是否等待并安装 replicate 的 unsharded Parameter这些 flag 不决定参数属于哪类,也不直接决定梯度通信。参数是否切分由
_init_hsdp_params()的enable_fsdp_shard决定;梯度通信由post_backward()的 compat/replicate/group 路径决定。5.2 [Hyper 当前]
init_hsdp_paramssequenceDiagram participant API as fully_shard API participant SCH as Scheduler participant ST as HSDPState participant P as HSDPParam participant M as Module API->>API: has_dtensor_param = any(...) API->>SCH: mesh 或 DTensor compat mesh SCH->>SCH: 构造 DDPMeshInfo/FSDPMeshInfo/HSDPMeshInfo SCH->>ST: new state(mesh_info) loop 每个 managed param ST->>ST: infer LOCAL/COMPAT/UNIFIED ST->>P: new HSDPParam(mesh_info, enable_fsdp_shard) P->>P: torch.chunk;所有 shard dim 要求整除 P->>P: clone actual shard -> _sharded_param_data P->>P: 由最终 layout 构造 sharded/unsharded GroupInfo P->>M: 原参数替换为 sharded_param ST->>ST: 分入 shard 或 replicate 列表 endreplicate_params使用shard_world_size=1创建与原 local 参数同 shape 的sharded_param。它不是 FSDP shard,只是通用状态机中的持久 Parameter。5.3 [Hyper 当前] unshard 细粒度顺序
以
reshard_after_forward=True为例:sequenceDiagram participant H as forward/backward pre-hook participant ST as HSDPState participant P as HSDPParam participant AG as shard ProcessGroup participant O as all_gather_output participant M as Module H->>ST: unshard(async_op, unshard_replicate) alt unshard_replicate 且 is_replicate_shard ST->>P: replicate_param.unshard() P->>P: all_gather_inputs(可 param_dtype cast) P->>O: allocate/resize full-size output P->>O: shard_size=1,完整 copy,无 collective end alt is_shard ST->>P: sharded_param.unshard() 或 param_group.unshard() P->>P: all_gather_inputs(可 param_dtype cast) P->>O: allocate/resize W 倍 output P->>AG: async AllGather end ST->>P: wait_for_unshard() P->>AG: handle.wait(如有) P->>P: unpack output,刷新稳定 _unsharded_param.data P->>M: module 参数替换为 _unsharded_param ST->>ST: 更新 is_shard/is_replicate_shardshard_size == 1和replicate_params当前都不会发 AllGather collective,但仍创建all_gather_output并执行完整 copy。forward post-hook 调用shard(shard_replicate=False):真正 shard 参数切回,replicate 参数保留 unsharded,供 backward 直接复用。backward prefetch 同样传unshard_replicate=False。5.4 [Hyper 当前] reshard 细粒度顺序
sequenceDiagram participant H as forward/post-backward hook participant ST as HSDPState participant P as HSDPParam participant U as _unsharded_param participant S as sharded_param participant O as all_gather_output participant M as Module H->>ST: shard(shard_replicate) alt 真正 FSDP shard 参数且当前 UNSHARDED ST->>P: to_sharded() P->>M: module 参数替换为 sharded_param P->>O: storage.resize_(0) end alt shard_replicate 且 replicate 当前 UNSHARDED ST->>P: to_sharded() P->>S: 当前实现:copy unsharded data -> same-shape sharded storage P->>M: module 参数替换为 sharded_param P->>O: storage.resize_(0) end Note over U,S: 该copy源于replicate参数维护两份storage,不属于目标设计这次 copy 只是当前实现为两份同 shape storage 补充的回写:如果 forward 原地修改了 unsharded 参数,它试图在切回对象前把值同步到持久 storage。真正 FSDP shard 参数不需要该 copy,因为 sharded storage 始终是 master owner,unsharded AllGather output 只是计算视图。目标方案不支持依赖 forward 原地改参的语义,因此该 copy 没有保留价值;它还可能在 mixed precision 下把 cast 后的低精度计算副本错误覆盖回 master。归一后 no-cast 路径直接 alias sharded storage,cast 路径释放计算副本且禁止 copy-back,两条路径都只切换 Parameter 映射。
5.5 目标决策:将
mesh_info下沉到HSDPParamMeshInfo的所有权从 state 下沉到 param,不等于每个参数都要深拷贝一份对象:普通 shard 参数若拓扑相同可以共享同一个不可变FSDPMeshInfo/HSDPMeshInfo引用;但 state、param group 和 executor 不得再假设一个 unit 内所有参数的 DP 通信拓扑都相同。HSDPParam.mesh_info:只描述 FSDP/HSDP 数据并行拓扑,包括 shard rank/size/group 和 replicate rank/size/group;HSDPParam保存的原 DTensor layout:只描述 TP/EP 等模型并行拓扑、placements 和 logical tensor meta;参数场景对应的
MeshInfo:MeshInfoFSDPMeshInfoHSDPMeshInfoHSDPMeshInforeplicate_paramsDDPMeshInforeplicate_paramsR*SDP ranks 的DDPMeshInfo约束:
HSDPState不再保存self.mesh_info;它只遍历参数并提交生命周期或梯度操作。mesh_info中的 shard/replicate group。None或 size 1 时,对应 FSDP/HSDP collective 是 identity/no-op。MeshInfo,也不由MeshInfo推导;第 6.4 节由HSDPParam从初始化时保存的原 Parameter DTensor mesh/placements 提取 replicate 轴并缓存 group,根反向钩子统一发起 TP AR。MeshInfo不保存 gradient、dtype、buffer、Work、Event 或 TP placement 等动态/模型并行状态。5.6 目标
init_hsdp_paramssequenceDiagram participant API as fully_shard API participant ST as HSDPState participant P as HSDPParam participant M as Module API->>API: 一次性校验全 Tensor 或全 DTensor API->>API: 校验显式 mesh 与原 layout loop 每个 hsdp param API->>API: 为参数选择 FSDP/HSDP/DDP MeshInfo API->>P: 原 Parameter、logical layout、param mesh_info P->>P: 计算 shard rank/size 与 actual/padded shape P->>P: 建立 sharded owner storage + actual view P->>P: 保存稳定 original/unsharded Parameter 对象 P->>M: 安装 sharded_param 持久态 ST->>ST: 追加到唯一 hsdp-param 列表 end Note over ST: 单一 ShardedState;无 replicate 专用列表/状态参数场景只改变事实,不改变 API:
shard_world_sizeneeds_all_gathermesh_info> 1FSDPMeshInfo/HSDPMeshInfo1replicate_params1DDPMeshInfo唯一判定为:
不保留
uses_param_shard:它同时重复了MeshInfo拓扑和shard_world_size通信规模,容易与其他状态 flag 组合膨胀。是否在 logical layout 上安装 FSDPShardplacement,从参数mesh_info是否具有shard_mesh_dim推导;是否执行 AllGather/ReduceScatter,只看shard_world_size;参数当前处于 sharded 还是 unsharded,则只由ShardedState表达。这样 shard size 1 仍保留 FSDP logical layout,而replicate_params仍是DDPMeshInfo,两者不会因 world size 都为 1 而混淆。5.7 目标 unshard:通信路径与本地 alias 路径归一
sequenceDiagram participant H as pre-hook/prefetch participant ST as HSDPState participant P as HSDPParam participant PG as P.mesh_info.shard_process_group participant B as owned temp/output storage participant U as stable unsharded Parameter participant M as Module H->>ST: unshard(async_op) loop 每个 managed param ST->>P: unshard(async_op) alt needs_all_gather P->>P: 读取 padded sharded input P->>P: 可选 cast 到 param_dtype P->>B: 分配/复用 AllGather output P->>PG: AllGather P->>PG: wait/event dependency P->>B: unpack + narrow 到 logical full shape P->>U: rebind data/view 到 AllGather owner else 无 AllGather,且不需要 param_dtype cast P->>U: rebind data/view 到 sharded/local owner 的稳定 view Note over P,U: data_ptr/storage 相同;不分配 all_gather_output;不 copy else 无 AllGather,但需要 param_dtype cast P->>B: cast sharded/local owner -> param_dtype temp P->>U: rebind data/view 到 cast 结果 Note over P,U: 不创建 AllGather output;U 引用 cast owner end P->>M: 安装同一个 stable unsharded Parameter 对象 end对
replicate_params和shard_size == 1:unsharded_param必须直接引用sharded_paramlocal storage 的 view。unsharded_param必须引用 cast 结果;融合路径中可引用共享 flat cast buffer 的对应 slice。5.8 目标 reshard 与 storage ownership
sequenceDiagram participant H as forward/post-backward hook participant ST as HSDPState participant P as HSDPParam participant M as Module participant AG as AllGather owner participant C as Cast owner participant S as Sharded owner H->>ST: reshard() loop 每个 managed param ST->>P: to_sharded() P->>M: 安装 stable sharded Parameter alt unsharded view 由 AllGather output 持有 P->>AG: 在最后一个 consumer/event 后 resize_(0) 或回收 else unsharded view 由 cast temp 持有 P->>C: 在最后一个 consumer/event 后释放引用/复用 buffer else unsharded view 借用 sharded owner P->>S: no-op;禁止 resize_(0),禁止 copy-back end P->>P: 状态置 SHARDED end建议显式记录 storage ownership,而不是从
is_sharded猜测:SHARDED_STORAGECAST_STORAGEparam_dtypeALL_GATHER_STORAGEFUSED_BUFFER_STORAGEoptimizer 只更新 sharded owner。mixed-precision cast view 是计算副本,不在 reshard 时反向覆盖 master 参数;forward 内原地修改 mixed-precision 参数不属于本期支持语义,必须通过用例或显式报错固定,不能隐式把低精度值 copy 回 master。
完成该归一后,删除
unshard_replicate、shard_replicate、wait_for_replicate、is_replicate_shard,并让 prefetch 遍历同一 managed-param 列表;无通信参数自然是幂等 no-op/alias。5.9
ReduceScatterPlan的当前作用与目标扩展[Hyper 当前]
platform/*/fully_shard/pack_utils.py::ReduceScatterPlan不是通信 group 或调度计划。它只描述一个参数在“module/grad 的原始 local full layout”与“集合通信要求的 row-major packed layout”之间如何转换:pack_kindidentity_dim0、same_dim_strided_identity_dim0或chunk_cat_non_dim0shard_dim、world_sizeunpacked_shapepacked_tensor_shapepacked_shape(world_size, per_rank_numel)viewpack_for_reduce_scatter()unpack_from_all_gather()当前 plan 还承担 AllGather output 的逆 pack,所以名称虽然是
ReduceScatterPlan,职责实际是“单参数 AG/RS layout plan”。它不选择 ProcessGroup、reduce op、dtype、stream、bucket 或 Work 生命周期。当前实现对所有 uneven 拒绝,并要求输入 contiguous。[目标] 保留这个边界。
ReduceScatterPlan不选择MeshInfo、TP replicate group 或 ProcessGroup,只补充或可推导以下信息:unpacked_shapeactual_sharded_shapepadded_sharded_shapepadded_unsharded_shapepacked_shape(world_size, padded_sharded_numel)actual_sharded_numelpadded_sharded_numelpack_kind均匀 dim-0 的
actual == padded,继续走无额外 padding 的 view 快路径。非 dim-0 uneven 在build_rs_plan()前的初始化校验中拒绝。5.10 去除HSDPState级 dtype强制要求一致的约束
当前HSDPState在初始化流程会会有一个初始化state级别的param_dtype, reduce_dtype的流程。并且要求所有HSDPParam的相关混合精度配置要完全一致。
对于comm_fusion=True的路径来说,当前保持该约束。
对于comm_fusion=False的路径来说,当前可以去掉这个约束。因为通信行为的粒度更小。此外MindFormer有场景是网络中初始化后参数的dtype较就混合着FP32与BF16。不原生支持这种方式的话,一个TransformerLayer要包7个fully_shard。造成比较大的host开销。
6. 反向梯度通信重构
6.1 重构前后对比
unsharded_param.grad、兼容路径下的sharded_param.grad、累积梯度unsharded_param.grad或其累积缓冲区取得普通 Tensorparam_mode、状态对象的mesh_info、GroupInfo、重复参数标记和最终布局共同决定HSDPParam.mesh_info;TP 只读取原 Parameter DTensor 的placementsCommContext混合fully_shard调用树共享的根调度上下文持有队列、桶、句柄和最终收尾状态 RL场景中如果有一个进程内给多个模型包fully_shard,类级别变量可能会有问题。reduce_partial()/redistribute()兼容路径TP 域内需要参数梯度全归约的是配置为
Replicate的少量参数,主要是归一化层权重和偏置。大权重通常已经按 TP切分,不进入这一阶段。因此本期不把 TP AR 插入逐模块流水,也不让它改变现有 FSDP/HSDP 的通信重叠路径。
6.2 当前
root_backward_hook的收尾顺序与 TP 插入点Torch 当前
_root_backward_hook()先调用self._backward_hook(),保证当前单元遗漏的post_backward()得到补执行,然后在最终归约分支中依次处理:TP 通信只能放在第 6 步之后、根反向钩子返回之前:
def _root_backward_hook(self, force_reduce=False): self._backward_hook() if apply_final_reduce or force_reduce: # 保持当前 FSDP/HSDP 收尾逻辑及其顺序。 self._finalize_comm_fusion_reductions() self._launch_last_hsdp_allreduce() self.hsdp_state.reduce_scattered_params() TorchHSDPStateV2.delay_apply_reduce_grads(self.hsdp_state.device) self.hsdp_state.reduce_params() # 新增位置:上述 FSDP/HSDP/DP 通信全部完成后才进入 TP 阶段。 self._allreduce_tp_replicated_param_grads()HSDPParam,不能只遍历当前self.hsdp_state。placements中没有Replicate、参数被冻结、当前没有梯度或本轮关闭梯度同步时直接跳过。post_backward()中 RS/AR 的等待点、发起顺序和现有桶;纯 FSDP/HSDP 场景在新增位置是空操作。“FSDP/HSDP 已结束”指对应集合通信已完成,不表示可以提前释放或下沉其结果缓冲区。
当前 Torch 仍用
scheduler_state != BACKWARD控制是否进入最终归约分支。第 3.6 节要求的最终形态仍是根回调无条件执行空队列安全的收尾;TP 阶段随该无条件收尾执行。
目标路径:
RS(S)(R,S)RS(S) -> AR(R)(8,1)AR(8)replicate_paramsAR(flat S)replicate_paramsAR(flat R*S)ShardRS(S)ReplicateRS(S)AR(TP)ReplicateRS(S) -> AR(R)AR(TP)replicate_params+ TPReplicateAR(flat DP)AR(TP)TP 阶段接收已经完成 FSDP/HSDP/DP 归约的本地梯度。通信输入必须是普通 Tensor;若内部参数梯度以 DTensor
形式挂载,只能取得其实际本地张量参加通信,不能调用
reduce_partial()或redistribute()。TP AR 使用参数的self.reduce_op_type,不重复应用gradient_scaling_factor。TP 参数数量较少时可由根回调逐参数调用
hsdp_param.reduce_tp_grad(reduced_grad)。如果后续需要减少小通信发起次数,再按(tp_replicate_group, self.reduce_op_type, dtype, device)建立独立 TP 桶;无论是否融合,都不能写入普通HSDP 的
AllReduceParamGroup,也不能提前到post_backward()中。6.3 FSDP/HSDP 跨模块流水保持不变
非融合主路径继续保持“等待上一组 RS → 发当前 RS → 发上一组 AR”,用更早模块的反向计算掩盖通信:
当前队列与目标归一关系:
pre_reduce_scatter_params:仅需要 RS 的参数,下一反向钩子或根回调等待后应用。pre_all_reduce_groups:RS 已发起,下一反向钩子的launch_prev_allreduce()等待 RS 后发起对应 AR。pending_all_reduce_groups:AR 已发起,根回调等待后应用。pre_all_reduce_params、pre_direct_all_reduce_grads:目标删除;replicate_params与普通 HSDP 参数一起按各自ProcessGroup 进入统一的多组 AR 调度。
CommContext.pre_param_group/all_reduce_param_group:融合路径的同类流水,TP 阶段不写入这两个字段。一次
post_backward()支持多个 AR 组即可,不需要为普通 HSDP 和replicate_params再建立固定的两类调度结构。分组键至少包含 ProcessGroup、归约类型、数据类型和设备;每个组独立持有融合缓冲区与通信句柄。
launch_prev_allreduce()按稳定顺序遍历全部非空组,根回调等待全部句柄。普通 HSDP 的输入仍依赖 RS 完成,replicate_params仍跳过 RS;统一的是调度入口,不是二者的数据依赖。立即在
M_i.post_backward()等待 RS 再发 AR 会失去RS(M_i)与backward(M_i-1)的重叠。详细的流和事件方案见grad_comm_overlap.md;本次调整只把少量 TP 复制参数的 AR 移到根回调尾部,不改变FSDP/HSDP 流水。对于这批已确认主要为归一化层权重和偏置的小参数,本 RFC 采用该分析文档中的“根回调尾部同步”
方案;分析文档对大规模 TP 梯度尾部同步的性能风险仍然成立。
6.4 归约类型、缩放和梯度累积
FSDP RS、HSDP/DP AR 和根回调中的 TP AR 使用同一个参数级
self.reduce_op_type:AVG:纯 FSDP 在分片组平均;HSDP 在分片组和重复组各平均一次,最终除数为S × R。SUM:各 DP 通信阶段使用 SUM,不隐式平均。replicate_params的展平 DP AR 使用用户选择的 DP SUM/AVG 语义。self.reduce_op_type,不从梯度或Partial放置策略推导另一种归约类型。gradient_scaling_factor只在整条反向归约链的第一个实际集合通信前应用一次;根回调中的 TP AR 不再缩放。6.6 目标
post_backward()与根回调固定流程post_backward()只负责 FSDP/HSDP/DDP:根反向钩子负责最终收尾:
comm_fusion=True时普通 HSDP 参数仍走HSDPParamGroup.foreach_reduce();replicate_params不执行 RS,但其 AR 仍由同一次
post_backward()按自身 ProcessGroup 组织和发起。全部 FSDP/HSDP/DP 组在根回调完成后,再统一执行同一个 TP 尾部阶段。
7. 支持参数在 dim-0 非均匀切分
7.1 三层数据模型
必须区分:
对参与 FSDP 切分的 local full parameter 定义:
这里采用 PyTorch
torch.chunk/Shard._local_shard_size_and_offset()语义,不采用“前 remainder 个 rank 多一个”的 balanced split。两者在部分 shape 上不同,例如D0=10,W=4的 PyTorch chunk 为3,3,3,1。7.2 [PyTorch 2.9] 标杆行为
FSDPParam._init_sharded_param():_chunk_with_empty(),支持D0 < W;sharded_size记录 actual shape;padded_sharded_param_size取 rank 0 chunk shape;sharded_param._local_tensor是 padded storage 的 actual narrow view;DTensorSpec.tensor_meta仍描述 logical global tensor;padding 不进入 placement/tensor_meta。foreach_reduce():_get_dim0_padded_size()构造 RS input;fsdp.chunk_cat把尾部自动补零;sharded_size。reset_sharded_param():_apply后若 local tensor 是 actual shape,则重新构造 padded storage;_sharded_param_data;PyTorch 当前即使均匀也执行
new_zeros + copy。Hyper 选择只在 uneven 时创建 padding,以保留均匀快路径;收益是减少 allocation/copy,风险是 executor 必须正确处理两种 storage owner,测试矩阵不能只覆盖 uneven。7.4
_init_sharded_param()与reset_sharded_param()这里必须区分参数对象与通信存储区,二者不能再统称为“分片参数”:
self.sharded_paramnn.Parameter,也是优化器必须持有和更新的参数。它是 DTensor,其_local_tensor只表示当前进程的实际分片,实际第 0 维允许为 0,不包含补齐元素。sharded_param_init_sharded_param()中构造实际本地分片的局部变量;非均匀切分时最终改为补齐存储区上的narrow视图,再用于创建self.sharded_param。self.sharded_sizeself.padded_sharded_param_sizepadded_sharded_param_init_sharded_param()在非均匀切分时创建的全零补齐存储区。self._sharded_param_datasharded_param.view(-1);非均匀切分时指向padded_sharded_param.view(-1)。它不是模块参数,也不能交给优化器。当前实现中的:
self._sharded_param_data = sharded_param.view(-1)只适用于
self.sharded_size == self.padded_sharded_param_size的均匀切分。当前all_gather_inputs直接读取self._sharded_param_data,_get_unsharded_param_data()又把all_gather_inputs[0]作为全收集输入。因此非均匀切分不需要在每次通信前临时补齐,但初始化时必须让self._sharded_param_data指向补齐后的存储区,保证各进程的输入元素数相同。目标
_init_sharded_param()的核心存储关系如下。实际实现先按当前逻辑完成offload_to_cpu和pin_memory处理,再进入下面的存储分支,保证实际分片和补齐存储区位于同一设备:chunks = _chunk_with_empty(param_data, shard_world_size, dim=shard_dim) sharded_param = chunks[shard_rank].clone().contiguous() self.sharded_size = sharded_param.size() self.contiguous_sharded_stride = make_contiguous_strides_for(self.sharded_size) self.padded_sharded_param_size = chunks[0].size() length = sharded_param.size(shard_dim) if sharded_param.numel() > 0 else 0 if self.sharded_size == self.padded_sharded_param_size: # 均匀切分:实际分片本身就是通信存储区。 self._sharded_param_data = sharded_param.view(-1) else: # 非均匀切分:通信存储区统一补齐,参数只暴露实际前缀。 padded_sharded_param = sharded_param.new_zeros( self.padded_sharded_param_size ) if sharded_param.numel() > 0: padded_sharded_param.narrow( dim=shard_dim, start=0, length=length, ).copy_(sharded_param) self._sharded_param_data = padded_sharded_param.view(-1) sharded_param = padded_sharded_param.narrow( dim=shard_dim, start=0, length=length, ) self.sharded_param = nn.Parameter(self.to_sharded_dtensor(sharded_param)) self.sharded_param.requires_grad_(param.requires_grad) self._setattr_on_modules(self.sharded_param)self.to_sharded_dtensor(sharded_param)使用self._sharding_spec中显式保存的逻辑全局shape、stride、dtype,以及self._spmd_mesh和self._spmd_placements。其中sharded_param.size()只用于实际本地形状,不能用于反推逻辑全局形状。非均匀切分时,self.sharded_param._local_tensor与self._sharded_param_data共享底层存储区,但前者只覆盖实际前缀,后者覆盖包含补齐元素的完整通信存储区。
reset_sharded_param()在load_state_dict(assign=True)、元设备初始化或模块_apply后按以下顺序重建:new_param = self._resolve_reset_param()取得模块当前注册的参数,并要求它是 DTensor;local_tensor = new_param.to_local()取得检查点或_apply提供的实际本地视图。第 0 维长度为 0 是合法输入。self._sharding_spec校验new_param的显式逻辑全局shape、stride和dtype,用self._spmd_placements校验new_param.placements,用self.sharded_size校验local_tensor.size()。补齐形状只由self.padded_sharded_param_size决定。self.sharded_size == self.padded_sharded_param_size,令local_tensor = local_tensor.contiguous()、self._sharded_param_data = local_tensor.view(-1)和local_view = local_tensor.detach();三者共享同一个实际分片存储区。padded_local_tensor = local_tensor.new_zeros(self.padded_sharded_param_size),把local_tensor复制到padded_local_tensor的实际前缀,再令self._sharded_param_data = padded_local_tensor.view(-1),并令local_view为padded_local_tensor.narrow(dim=shard_dim, start=0, length=length).detach()。self.sharded_param._local_tensor = local_view,再执行self._sharding_spec = self.sharded_param.layout和self._setattr_on_modules(self.sharded_param)。不需要新增含义不清的storage_owner字段;self._sharded_param_data对完整通信存储区的强引用负责保持该存储区存活。local_tensor表示的实际分片,不保存补齐元素。非均匀分支必须每次通过new_zeros()重建补齐存储区,不能读取local_tensor实际范围之外的内存,也不能把旧补齐区中的内容当作已加载数据。
7.5 Unshard/AllGather
flowchart LR A[actual local shard view] --> P[padded communication storage] P --> AG[AllGather 等长输入] AG --> F[padded full buffer: C*W] F --> N[narrow/as_strided 到 logical D0] N --> U[module 使用的 unsharded Parameter]even 场景
actual == padded,A可直接作为通信输入;uneven 场景 padding 在参数 init/reset 时准备,不在每次 AllGather 前重复分配。若param_dtype需要 cast,cast 的对象是 padded communication input,保证每个 rank input numel 一致。7.6
reduce_scatter_grad()unpacked_shape。padded_unsharded_shape,尾部[D0, C*W)必须置 0。padded_sharded_numel。actual_sharded_shape建 grad view,padding 不挂到 optimizer grad。global dim-0 < W时,后部 rank actual grad 为(0, *rest),但 collective input/output 仍使用统一C。所有numel==0分支必须保留非 shard 维 shape,不能退化为无维度 empty tensor。7.7 示例
示例 A:global
(5,3),FSDP world size 2(3,3)(3,3)(2,3)(3,3)示例 B:global
(2,3),FSDP world size 4(1,3)(1,3)(1,3)(1,3)(0,3)(1,3)(0,3)(1,3)示例 C:TP-local dim-0 再被 FSDP uneven
global
(10,2),root mesh(dp=2,tp=2),TP 与 FSDP 都切 dim-0。每个 TP-localD0=5,再按 FSDP 切成 3 和 2:(dp,tp)(0,0)(3,2)(3,2)(0,1)(3,2)(3,2)(1,0)(2,2)(3,2)(1,1)(2,2)(3,2)最终 placement 为
(_StridedShard(dim=0, split_factor=2), Shard(0))的等价 HyperStridedShard表达。split_factor和 offset 只描述 logical shard 顺序,不描述 padding。PyTorch 2.9 的 shape/offset helper限制
_StridedShard段结束后不能再对同一 tensor dim 继续 sharding;Hyper 应复用同等校验。二维 TP+FSDP 可支持,更高维连续同 dim sharding必须有独立用例,不能从二维结果外推。示例 D:均匀快路径
global
(8,3)、W=2:两 rank actual/padded 都是(4,3)。不得创建独立 padding storage,不增加 zero/copy。7.8 DCP 与 state dict
[PyTorch 2.9] DCP 使用 DTensor logical shape、placements 和
compute_local_shape_and_global_offset()生成ChunkStorageMetadata,__get_tensor_shard__()返回 actualto_local()。load planner按 saved chunk 与目标 actual chunk 的交集生成 read items;FSDP load post-hook 再重建 padded storage。[Hyper 当前]:
StandardSavePlanner和create_chunk_list_for_tensor()使用distributed_checkpoint/reshard.py::infer_slice_area_by_rank();该函数对任何非整除直接抛ValueError。core/dtensor/layout.py::_infer_slice_area_by_rank()则使用 floor,可能静默丢 remainder。core/utils/shape_utils.py并不一致。[目标]:
8. 验证设计
8.1 用例分层
8.2 Hook 状态机用例
保留第 3.7 节已有三种场景,并新增:
requires_grad=True,验证PostBackwardFunction先进入BACKWARD后 root 仍 drain。force_reduce与自然 root callback 重复到达时不重复应用 grad。8.3 支持参数在 dim-0 非均匀切分的测试用例
本小节是 uneven 功能的独立验收集,不能被普通 FSDP 精度用例替代。
8.3.1 Shape、offset 与 logical meta UT
(5,3), W=2(3,3)/(2,3);offset0/3;两 rank logical global shape 都是(5,3)(2,3), W=4(1,3)/(1,3)/(0,3)/(0,3);empty offset 为 2(10,3), W=43/3/3/1,防止误实现 balanced3/3/2/2(8,3), W=2(0,N)行为有固定契约8.3.2 参数 storage 与生命周期
_init_sharded_param():actual view/padded owner/data_ptr/尾部 0。reset_sharded_param():load actual view 后重新 padding,重复调用幂等。param_dtype=None:no-AG路径的 unsharded/sharded storage data_ptr相同。param_dtype!=orig_dtype:no-AG路径引用 cast结果,不创建 AllGather output;reshard后master未被低精度覆盖。replicate_params和 shard size 1:无 AllGather allocation/copy;2-D HSDP replicate group为 flattenedR*S。8.3.3 TP 组合
D0 < W。StridedShard、split factor、actual shape、global offset和DCP metadata。8.3.4 训练结果
main_grad。所有新增 uneven 用例在实现前均标记为“待实现/待执行”,不得写成已通过。
8.4 DCP/state dict 用例
至少覆盖:
_StridedShardsave/load。assign=True、meta/deferred init。_sharded_param_datapadded size正确、尾部全0,checkpoint中无 padding payload。现有均匀 DCP reshard用例可复用框架,但不能替代上述 uneven断言。
8.5 通信与混合并行矩阵
(8,1)、replicate_params、shard size 1每个参数的 debug view 要与实际 profiler collective group、op 和 payload 对上。
TP 路径至少单独断言:gradient 始终是普通 Tensor;原 Parameter DTensor 无 replicate placement 时根反向钩子不发 TP 通信;有 replicate placement 时在对应 group 上以
self.reduce_op_type发起 AR;多 replicate 轴按约定顺序处理;TP AR 只在全部 FSDP/HSDP/DP 通信完成后执行,且不重复应用 gradient scaling。纯 FSDP/HSDP 用例必须断言 TP 阶段为空操作,原有 RS/HSDP AR 的桶数、发起顺序和负载不变。多通信组至少单独断言:普通 HSDP 与
replicate_params同时存在且使用不同 ProcessGroup 时,同一次post_backward()能完成分组,launch_prev_allreduce()对每个非空组各发起一次 AR 并分别保存句柄;根回调等待全部组后才允许发起 TP AR。
8.6 性能与显存验收
每次性能记录必须包含:
初始阻断门槛:已有均匀场景 throughput 不低于 baseline 97%,且无稳定回退趋势;峰值显存不高于 baseline。最终门槛应根据目标NPU环境方差收紧。均匀场景不得出现为uneven新增的padding allocation/copy;shard size 1和replicate路径不得出现AllGather或同shape output copy。
MindFormers E2E至少覆盖 FSDP、HSDP、TP+FSDP、TP+HSDP、recompute组合、PP组合和checkpoint续训。无法执行的多卡用例必须保留准确命令和“待执行”状态,不填造性能数据。
9. 实现拆分与代码改动点
9.1 建议 PR 顺序
MeshInfo从state下沉到param;按参数分配FSDP/HSDP mesh_info;保存原DTensor TP replicate group每个 PR 必须可运行。迁移适配层应局限在一个模块并在对应 PR 内删除,不能先删旧路径、后续 PR 才恢复正确性。
9.2 模块改动
core/fully_shard/api.pycore/fully_shard/hsdp_scheduler.pycore/fully_shard/hsdp_state.pycore/fully_shard/hsdp_param.pyplatform/torch/fully_shard/param.pyplatform/torch/fully_shard/pack_utils.pyReduceScatterPlandim-0 padding与actual/padded信息;非dim-0 uneven拒绝platform/torch/fully_shard/param_group.pyplatform/torch/fully_shard/state.pylaunch_prev_allreduce()发起;删除compat/direct旁路platform/torch/fully_shard/scheduler.pycore/dtensor/*core/distributed_checkpoint/*platform/mindspore/fully_shard/*10. 兼容性与方案取舍
mesh=None + DTensorMeshInfo所有权HSDPParam,只管FSDP/HSDP公开
fully_shard()签名不变。行为变化包括:混合参数和mesh=None + DTensor从深层/隐式行为变为初始化早期错误;这是有意收紧。11. 风险与待确认项
_StridedShard三维以上同dim顺序resize_(0)reduce_op_type重复缩放已确定、无需重新讨论的方向:Hook无条件root drain、单次类型拦截、删除执行mode、MeshInfo下沉到参数且只管FSDP/HSDP、梯度始终是普通Tensor、单次post_backward支持按实际ProcessGroup处理多个AR组、TP通信只在root尾部按原Parameter DTensor replicate placements触发、no-AG alias、只支持dim-0 uneven、按需padding、logical/actual/padded分离、DCP不保存padding。
12. Definition of Done