已开启
fsdp优化 #184
changzherui创建于 6月3日
6月12日 修改了issue 的描述
6月12日 修改了issue 的描述
6月12日 将 changzherui1 设为负责人
6月12日 修改了issue 的描述
6月13日 修改了issue 的描述
6月13日 修改了issue 的描述
6月15日 关联了pull request:fix(mindspore): align replicate_params drain with Torch and add TP+FSDP ST
6月15日 关联了pull request:fix(mindspore): align replicate_params drain with Torch and add TP+FSDP ST
Hyper通信掩盖方案
Torch FSDP2 在分布式参数分片通信流程中,会先对分片参数进行合并处理,再执行通信操作,这一机制会引入显著的 copy_in 与 copy_out 数据拷贝开销。为此,HyperParallel 针对性优化通信逻辑:在权重 Unshard 的 AllGather 过程、以及梯度同步的 ReduceScatter 过程中,统一使用逐权重下发模式,跳过参数合并与批量拷贝环节,大幅削减冗余数据搬运,优化分布式训练的整体性能。
前向传播阶段采用分层预取与即时重分片相结合的策略实现计算与通信重叠:模型每一层执行前,会通过 forward_pre_hook 完成当前层参数的 unshard 下发和通信同步操作,同时异步发起下一层参数的 unshard 通信预取,提前加载后续层权重数据;当单层前向计算完成后,再借助 forward_hook 及时执行权重 reshard 重分片操作,快速释放当前层占用的显存资源。该机制利用 Hook 回调的时序特性,将参数分片通信、权重加载与层计算任务并行执行,有效掩盖通信延迟,同时动态管控显存占用,兼顾分布式训练的通信效率与显存利用率。
反向传播阶段的通信掩盖逻辑更为复杂,流程同时覆盖参数unshard通信、梯度两类不同域的同步操作,包含Shard 域 ReduceScatter与Replicate 域 AllReduce。方案依托反向传播的 Hook 回调链路完成多段通信的分层调度:在backward_pre_hook中提前异步发起下一层参数的unshard通信预取;进入backward_hook后,按时序依次执行三步操作:首先完成上一层 Shard 域的梯度 ReduceScatter 同步,接着异步提交当前层 Shard 域的 ReduceScatter 通信,最后以融合方式批量下发上一层梯度在 Replicate 域的 AllReduce 通信,且在backward末尾的root module再执行Replicate域AllReduce通信的同步wait,让它能和反向计算充分掩盖。由于两类梯度同步分属不同通信域,彼此可完全并发执行,进一步挖掘并行潜力。整套调度策略将参数解分片、两级梯度同步等多类通信任务与反向计算流程深度穿插,最终实现三重通信掩盖,大幅削弱分布式场景下的通信时延损耗,整体通信效率得到显著提升。
考虑到融合 AllReduce 通信要求输入数据必须位于连续内存空间,我们针对性设计了零拷贝内存优化:预先分配一块连续的内存缓冲区,将 ReduceScatter 的输出结果直接写入该缓冲区。各权重梯度经过 ReduceScatter 计算后,数据天然在缓冲区中连续排布,无需额外数据拷贝,便可直接基于这块连续内存执行融合 AllReduce 操作,彻底规避数据搬运带来的性能开销。
前向传播依靠分层预取unshard与即时reshard,让参数加载通信和层计算并行执行;反向传播则通过多阶通信调度、跨域通信并发、梯度通信融合以及零拷贝优化,实现三重通信掩盖。整套方案将各类分片、梯度同步通信深度穿插在计算流程中,最大化利用硬件空闲资源,有效掩盖分布式训练中的通信延迟,同时兼顾显存管控与数据传输效率,显著提升整体训练吞吐。
方案落地说明(PR #758:MindSpore FSDP2 comm_fusion 对齐 Torch 后端)
本评论记录上述「Hyper 通信掩盖方案」在 MindSpore 后端的具体实现、达成效果,以及与改动前的区别。
一、本 PR 做了什么
fully_shard的通信融合路径对齐 Torch 后端,在 Ascend 上落地「前向 AllGather 预取 + 即时 reshard」与「反向三重通信掩盖」。ag_output)显存泄漏。二、具体怎么做的
param_group.py)_pack_shards_into_ag_input_slice把各权重 local shard 直接打包进ag_output的本 rank slice(FSDP 融合 AG 实现约束:send buffer 必须是ag_output本 rank slice,与all_gather_copy_in()一致;见下文「send buffer 布局澄清」),省去all_gather_inputs的 cast 临时量;copy-out 用foreach_copy_without_bumping_version批量拷贝。reduce_scatter_copy_in用一次mint.cat打包,替代逐 rownarrow+copy。AllReduceParamGroup:把同一 replicate group 内多个 bucket 聚合成一次融合 all-reduce;分配 512B 对齐的连续fused_buffer,ReduceScatter 输出直接写入 buffer view(零拷贝),再基于这块连续内存执行融合 AllReduce。state.py+scheduler.py)delay_apply_reduce_grads统一 wait,让 AllReduce 与反向计算充分重叠。reduce_params/reduce_scattered_params拆分,职责与 Torch 后端对齐。_resolve_default_reduce_op:managed 参数含 DTensor →SUM,否则AVG。_need_div/_div_if_needed全部手动除法分支;融合 all-reduce 因 buffer 含对齐 padding 而用SUM,随后在wait_and_apply_grads中/replicate_world_size得到等价 AVG。replicate_world_size正确缩放,与融合路径数值一致。ag_output泄漏:foreach_all_gather_copy_out在 copy-out 后调用free_all_gather_output()(resize_(0))释放融合 buffer storage(per-paramall_gather_outputs已持有数据),保留 tensor 对象供下轮复用。三、达到的效果
(2,4)+ SGD 下,comm_fusion=False与comm_fusion=True逐步 loss 逐位一致;UT 49 项全过,8 卡 Ascend910B ST(精度/对齐/grad accum)通过。四、与改动前的区别
SUM+ 手动_div_if_needed,replicate 维漏除导致梯度偏大AVG/SUM;融合路径用SUM+/replicate_world_size_flat_param_bufferrebase + 独立 buffer 作 send(FSDP 路径曾触发HcclAllGather ret:2)_pack_shards_into_ag_input_slice打包进ag_output本 rank slice(与 copy_in 同布局)ag_outputresize_(0)释放 storageAllReduceParamGroup融合 + 跨域(Shard/Replicate)并发reduce_params混合 RS/ARreduce_params与reduce_scattered_params拆分,对齐 Torch推荐用法:
fully_shard(..., comm_fusion=True, comm_fusion_zero_copy=False)。四.1 MindSpore 融合 AllGather:send buffer 布局澄清(修正 HCCL 表述)
此前文档/注释曾写成「HCCL
all_gather_into_tensor一律要求 send buffer 必须是 output 本 rank slice,NCCL 允许分离」——该表述过绝对,此处更正。依据来源(本仓库内,非 HCCL 官方手册摘录):
comm_fusionzero_copy 路径若用独立_flat_param_buffer或 storage rebase 仿 Torch 作 send buffer,在 融合 AG + async prefetch +ag_outputstorage resize/free 组合下会出现RuntimeError: HcclAllGather failed, ret:2(见test/qwen3_4b_fsdp2/docs/MS_ZERO_COPY_HCCL_ROOT_CAUSE.md)。test/qwen3_4b_fsdp2/diagnose_ms_zero_copy_ag.py,8 卡):input = out.narrow(rank*n, n)与 独立separate_intensor 均可 PASS——说明 MindSpore/HCCL API 层并非禁止独立 send buffer。all_gather_copy_in()布局一致、规避上述 FSDP 路径 ret:2,统一规定 fusion AG 的 send =ag_output.narrow(rank * input_numel, input_numel);comm_fusion_zero_copy=True时用_pack_shards_into_ag_input_slice批量 pack,而非 Torch 式_flat_param_bufferrebase。与 Torch/NCCL 的差异(准确说法):
_flat_param_bufferag_output本 rank slice(实现不变量)代码注释:
param_group.py中_pack_shards_into_ag_input_slicedocstring 已同步改为上述表述。先澄清一个事实前提:当前实现里
ag_output并不常驻。foreach_all_gather_copy_out在 copy-out 之后每一步都会调用free_all_gather_output()(param_group.py内storage.resize_(0)),把融合 buffer 的存储立即释放到 0 字节,只保留一个空 tensor 壳供下一轮alloc_all_gather_outputresize 回来复用。所以稳态下并没有「ag_output 常驻」这笔显存开销,谈不上「显存上已经付了常驻代价」。基于这一点,现状其实就是对齐 Torch 后端的内存/拷贝画像,而不是两头不沾的中间态:
comm_fusion_zero_copy=True走_pack_shards_into_ag_input_slice,直接把 shard 打包进ag_output的本 rank slice(轻量 pack)。注意:这不是 Torch flat-buffer 那种「完全跳过 copy-in」;MS 侧 zero_copy 语义是批量 pack 进 output slice,省 cast 临时量与 Python 循环。send 必须落在本 rank slice 是 本仓库 FSDP fusion 路径的已验证不变量,不是 HCCL 全局 API 禁止独立 input(裸 collective 探针见下)。foreach_all_gather_copy_out同样是torch.split_with_sizes_copy(ag_output, …, out=per-param all_gather_outputs),把融合结果拷进各参数独立的 unsharded 存储。free_all_gather_output()立即释放(Torchparam_group.py:691注释即「Immediately release fused buffer memory」)。也就是说,你说的 option 1(reshard 时补回
free_all_gather_output()、保留 pack)当前代码已经在做,而且更早——我们是在 copy-out 时就 free(早于 reshard),pack 也保留。所以 PR 现状 ≈ 你的 option 1,旧的内存语义并没有丢。你看到的可能是去掉 free 的某个中间版本?现在 free 就在foreach_all_gather_copy_out行尾。唯一的瞬时峰值是 copy-out 期间
ag_output(full) 与 per-paramall_gather_outputs(full) 短暂并存,随后ag_output立刻 free——这一点与 Torch 完全相同,不是本 PR 引入的额外代价。关于 option 2(把 unsharded param rebase 进
ag_output的 rank slice、连 copy-out 也省、AG 纯 in-place 的真零拷贝):它其实比 Torch 还激进——Torch 也没有把 unsharded param rebase 到 AG 输出上,它同样 copy-out。代价正是「优化器必须原地更新 storage」这个脆弱不变量,也正是旧_flat_param_buffer路线的痛点和不设默认的原因。在本 PR 采用的 send=output-slice 布局下,copy-out 的省略只能靠 unsharded param 直接 viewag_output,这就要求ag_output真常驻 + 原地 optimizer 兼容,权衡较大。结论:当前实现等同于「free + pack」即你的 option 1,与 Torch 内存语义一致,不存在 ag_output 常驻。option 2 的真零拷贝更适合作为独立的后续性能优化单独评估(它换来的是更脆弱的 storage 不变量),本 PR 聚焦正确性与 Torch 对齐。如果评审同意,我在 issue #184 里登记 option 2 作为 follow-up。
五、性能收益场景分析(Qwen3-8B HSDP benchmark 验证)
5.1 前提:
comm_fusion与 PR 的关系comm_fusionFalse(需显式fully_shard(..., comm_fusion=True))comm_fusion_zero_copyFalse(即使开了 fusion)PR 改动可分成 三条独立收益线:
_need_div/_div_if_needed手动除法(正确性为主)。comm_fusion=False:MindSpore 新增AllReduceParamGroup4-step 层间 RS/AR overlap(preopt MS 侧无此路径);reduce_params/reduce_scattered_params拆分对齐 Torch。comm_fusion=True:foreach_all_gather/foreach_reducefusion 路径;层间 RS→AR pipeline(comm_ctx);AG buffer 用后释放、foreach_copy、mint.catpack 等微优化。5.2 场景 1:最可能看到 backward 性能提升 —
comm_fusion=False+ HSDP旧路径(preopt,
comm_fusion=False):每层post_backward对每个参数 RS 发出后 立刻同步 wait,再 AR,几乎无层间 overlap。新路径(post):每层 4-step 流水线 — wait 上一层 RS → apply 纯 FSDP RS → 本层异步 RS(
AllReduceParamGroup)→ 上一层异步 AR;root backward 统一delay_apply_reduce_grads。受益条件(需同时满足):
comm_fusion=Falsefully_shard(block)(多 FSDP unit)当前 benchmark 未覆盖此路径(实验均显式
--comm-fusion)。5.3 场景 2:
comm_fusion=True— backward 有小优化,forward 可能拖后腿Backward 侧改进:
ReduceOp.AVG/SUM,去掉 fusion 路径 RS/AR 后的div_()(preoptneeds_avg_div)reduce_scatter_copy_in用mint.cat+ 批量 copyForward 侧可能变慢:
comm_fusion_zero_copy=True时生效;默认 fallback 到reset_sharded_param+all_gather_copy_ininit_unsharded_param增加 contiguous 检查/拷贝Profiler(prefetch=1 + fusion + profile_stages,32 层 seq=512):
End-to-end step time(18 步均值,warmup=2):
zero_copy 对 pre 约 +6.2%,对 post ≈0% → post 的 AG pack 路径未兑现;整体 forward 退化 > backward 改进。
5.4 场景 3:主要是正确性/数值,不是性能
ReduceOp.SUM:修复 preopt AVG 映射 + 手动除法的语义问题AllReduceParamGroup用 SUM AR + 按需/ replicate_world_size:HSDP replicate 维平均化正确run_fully_shard_hsdp_avg_grad_scale_parity:fusion / 非 fusion 在 SGD 下 loss 一致对 Adam + DTensor 大模型,这更多是训练正确性,不是吞吐提升来源。
5.5 场景 4:内存收益(间接性能)
foreach_all_gather_copy_out后free_all_gather_output(),降低 AG buffer 峰值算力不受限时 step time 可能不变,但能跑更大配置。
5.6 场景 5:基本看不到提升
fully_shard(model):层间 pipeline 空间小comm_fusion=True且 forward 通信 bound(Qwen3 + prefetch + seq=512):实测 net 为负5.7 Qwen3 benchmark 配置与 PR 优化块对应
结论:在 Qwen3 +
comm_fusion=True场景下,PR 的 net step time 为负,主要因 fusion forward AllGather 路径退化;PR 设计上的主要性能红利在comm_fusion=False的 HSDP 层间 RS/AR overlap,当前 benchmark 未测到。5.8 建议 follow-up 验证
comm_fusion=False+ HSDP + prefetchcomm_fusion=True,comm_fusion_zero_copy=True,对比 pre/post