mindspore-ai / mindspore-ai/hyper-parallel
fsdp优化
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 53
- Forks
- 63
- Avg merge
- 23h 45m
- Merged PRs (30d)
- 63
Description
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 做了什么
- 将 MindSpore
fully_shard的通信融合路径对齐 Torch 后端,在 Ascend 上落地「前向 AllGather 预取 + 即时 reshard」与「反向三重通信掩盖」。 - 顺带修复 review 阶段发现的两个正确性问题:HSDP 非融合路径的 AVG 梯度缩放错误、融合 all-gather buffer(
ag_output)显存泄漏。
二、具体怎么做的
- 融合通信打包(
param_group.py)- fused AllGather:
_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批量拷贝。 - fused ReduceScatter:
reduce_scatter_copy_in用一次mint.cat打包,替代逐 rownarrow+copy。 - 新增
AllReduceParamGroup:把同一 replicate group 内多个 bucket 聚合成一次融合 all-reduce;分配 512B 对齐的连续fused_buffer,ReduceScatter 输出直接写入 buffer view(零拷贝),再基于这块连续内存执行融合 AllReduce。
- fused AllGather:
- 反向 4-step pipeline(
state.py+scheduler.py)- step1:wait 上一层 Shard 域 ReduceScatter;step2:apply 纯 FSDP 参数的 RS 结果;step3:异步发起当前层 RS(按 replicate group 融合分组);step4:融合批量下发上一层 Replicate 域 AllReduce;在 root module 末尾用
delay_apply_reduce_grads统一 wait,让 AllReduce 与反向计算充分重叠。 - 将
reduce_params/reduce_scattered_params拆分,职责与 Torch 后端对齐。
- step1:wait 上一层 Shard 域 ReduceScatter;step2:apply 纯 FSDP 参数的 RS 结果;step3:异步发起当前层 RS(按 replicate group 融合分组);step4:融合批量下发上一层 Replicate 域 AllReduce;在 root module 末尾用
- 梯度缩放统一为原生 HCCL AVG/SUM
_resolve_default_reduce_op:managed 参数含 DTensor →SUM,否则AVG。- 删除 legacy
_need_div/_div_if_needed全部手动除法分支;融合 all-reduce 因 buffer 含对齐 padding 而用SUM,随后在wait_and_apply_grads中/replicate_world_size得到等价 AVG。
- 两个正确性修复
- AVG 缩放:非融合 2D HSDP 的 replicate 维 all-reduce 现按
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 对象供下轮复用。
- AVG 缩放:非融合 2D HSDP 的 replicate 维 all-reduce 现按
三、达到的效果
- 通信掩盖:参数 unshard、Shard 域 ReduceScatter、Replicate 域 AllReduce 三类通信与计算深度穿插并发,HSDP 下 RS/AR overlap 明显改善。
- 显存:融合 all-gather buffer 用后立即释放,消除常驻显存泄漏,保住 FSDP 的显存收益。
- 数值正确性:2D HSDP
(2,4)+ SGD 下,comm_fusion=False与comm_fusion=True逐步 loss 逐位一致;UT 49 项全过,8 卡 Ascend910B ST(精度/对齐/grad accum)通过。
四、与改动前的区别
| 维度 | 改动前 | 改动后 |
|---|---|---|
| AVG 缩放 | SUM + 手动 _div_if_needed,replicate 维漏除导致梯度偏大 |
原生 HCCL AVG/SUM;融合路径用 SUM + /replicate_world_size |
| copy-in 零拷贝 | _flat_param_buffer rebase + 独立 buffer 作 send(FSDP 路径曾触发 HcclAllGather ret:2) |
_pack_shards_into_ag_input_slice 打包进 ag_output 本 rank slice(与 copy_in 同布局) |
ag_output |
copy-out 后未释放,常驻显存(泄漏) | copy-out 后 resize_(0) 释放 storage |
| 反向调度 | 单段 reduce | 4-step pipeline + AllReduceParamGroup 融合 + 跨域(Shard/Replicate)并发 |
| reduce 入口 | reduce_params 混合 RS/AR |
reduce_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 官方手册摘录):
- 调试现象:旧 MS
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)。 - 裸 collective 探针(
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。 - PR 实现选择:为与
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 的差异(准确说法):
| Torch(NCCL) | MindSpore FSDP fusion(本 PR) | |
|---|---|---|
| send buffer | 可为独立 _flat_param_buffer |
必须为 ag_output 本 rank slice(实现不变量) |
| zero_copy 含义 | rebase 后 AG 直读 flat buffer | 每步 pack shard 进 output slice(仍有 pack copy) |
| 根因 | NCCL 接受独立 input | FSDP 内存生命周期 + 已观测 ret:2,非「HCCL 全局规则」 |
代码注释:param_group.py 中 _pack_shards_into_ag_input_slice docstring 已同步改为上述表述。
先澄清一个事实前提:当前实现里 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_output resize 回来复用。所以稳态下并没有「ag_output 常驻」这笔显存开销,谈不上「显存上已经付了常驻代价」。
基于这一点,现状其实就是对齐 Torch 后端的内存/拷贝画像,而不是两头不沾的中间态:
- copy-in:
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 探针见下)。 - copy-out:两个后端都拷贝。Torch 的
foreach_all_gather_copy_out同样是torch.split_with_sizes_copy(ag_output, …, out=per-param all_gather_outputs),把融合结果拷进各参数独立的 unsharded 存储。 - 释放融合 buffer:两个后端都在 copy-out 末尾
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-param all_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 直接 view ag_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 验证)
基于 PR #758(
61a0750)代码审阅 + Qwen3-8B-lite(32 层、seq=512、HSDP mesh=(2,4)、8 卡 Ascend)pre/post 对比 benchmark。
pre = merge-basef993e54;post = PR 分支。两边均comm_fusion=True(除非另行说明)。
5.1 前提:comm_fusion 与 PR 的关系
| 开关 | 默认值 | 与 PR 关系 |
|---|---|---|
comm_fusion |
False(需显式 fully_shard(..., comm_fusion=True)) |
PR 同时改了 fusion / 非 fusion 两条 backward 路径 |
comm_fusion_zero_copy |
MindSpore 默认 False(即使开了 fusion) |
PR 把 zero-copy 从「持久 flat buffer rebase」改成「AG 时 pack 进 rank slice」 |
PR 改动可分成 三条独立收益线:
- 始终生效:reduce op 统一(DTensor→SUM,否则 AVG);移除
_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=False- 2D HSDP mesh(replicate > 1,存在 RS + AR 两阶段)
- 逐层
fully_shard(block)(多 FSDP unit) - 单层 backward 计算时间 ≥ 单层 RS/AR 通信时间
当前 benchmark 未覆盖此路径(实验均显式 --comm-fusion)。
5.3 场景 2:comm_fusion=True — backward 有小优化,forward 可能拖后腿
Backward 侧改进:
- 原生 HCCL
ReduceOp.AVG/SUM,去掉 fusion 路径 RS/AR 后的div_()(preoptneeds_avg_div) reduce_scatter_copy_in用mint.cat+ 批量 copy- HSDP bucket AR 层间 pipeline(preopt 已有,PR 主要是数值语义对齐)
Forward 侧可能变慢:
- AG zero-copy pack 仅在
comm_fusion_zero_copy=True时生效;默认 fallback 到reset_sharded_param+all_gather_copy_in - 去掉 persistent flat buffer rebase(preopt init 时 rebase,post 每次 AG 再 pack/copy)
init_unsharded_param增加 contiguous 检查/拷贝
Profiler(prefetch=1 + fusion + profile_stages,32 层 seq=512):
| 阶段/算子 | pre → post | 变化 |
|---|---|---|
| forward(profile_stages) | 943 → 1327 ms | -40.8% |
| backward | 2067 → 1244 ms | +39.8% |
| allGather kernel(每 rank 均值) | 1421 → 1847 ms | -30% |
| ALLTOALL API | 14133 → 11236 us | +20.5% |
End-to-end step time(18 步均值,warmup=2):
| 配置 | pre ms/step | post ms/step | post vs pre |
|---|---|---|---|
| prefetch=1,no zero_copy | 1556.56 | 1624.80 | -4.4% |
| prefetch=1,zero_copy | 1459.88 | 1627.37 | -11.5% |
| 无 prefetch,no zero_copy | 1738.78 | 1795.32 | -3.3% |
zero_copy 对 pre 约 +6.2%,对 post ≈0% → post 的 AG pack 路径未兑现;整体 forward 退化 > backward 改进。
5.4 场景 3:主要是正确性/数值,不是性能
- DTensor 参数默认
ReduceOp.SUM:修复 preopt AVG 映射 + 手动除法的语义问题 AllReduceParamGroup用 SUM AR + 按需/ replicate_world_size:HSDP replicate 维平均化正确- ST
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 峰值- 长序列 / 大 batch / 多 prefetch 时,减少 OOM 或 cache thrashing
算力不受限时 step time 可能不变,但能跑更大配置。
5.6 场景 5:基本看不到提升
- 纯 1D FSDP(无 replicate AR):AllReduceParamGroup 收益有限
- 整模型单个
fully_shard(model):层间 pipeline 空间小 - 参数很少的小网络:fusion 打包 overhead > 通信节省
- 已开
comm_fusion=True且 forward 通信 bound(Qwen3 + prefetch + seq=512):实测 net 为负
5.7 Qwen3 benchmark 配置与 PR 优化块对应
HSDP (2,4) · 32 层 per-block fully_shard · comm_fusion=True · prefetch=1 · DTensor · MS zero_copy 默认 off
| PR 优化块 | 是否跑到 | 实测倾向 |
|---|---|---|
| AllReduceParamGroup 4-step overlap | 否(fusion 路径绕过) | 未验证 |
| fusion backward 优化 | 是 | backward 变快 |
| fusion forward AG 重构 | 是 | AllGather 变慢,主导 step |
| reduce op 统一 | 是 | 正确性为主 |
| zero_copy pack | pre 有效 / post 无效 | post 未兑现 |
结论:在 Qwen3 + comm_fusion=True 场景下,PR 的 net step time 为负,主要因 fusion forward AllGather 路径退化;PR 设计上的主要性能红利在 comm_fusion=False 的 HSDP 层间 RS/AR overlap,当前 benchmark 未测到。
5.8 建议 follow-up 验证
| 实验 | 目的 |
|---|---|
同样 Qwen3,comm_fusion=False + HSDP + prefetch |
验证 MS 新增 AllReduceParamGroup overlap |
comm_fusion=True,comm_fusion_zero_copy=True,对比 pre/post |
看 post 能否兑现 AG pack 收益 |
| 短 seq / backward-heavy(大 batch、小 seq) | forward 占比低时 backward 改进更易 net 为正 |
| Profiler 拆 forward AG:pack vs collective vs wait | 定位 post forward 退化根因 |
schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 184
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/184
Contributor guide
No contributing guide indexed for this repository
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start by reading param_group.py, state.py, and scheduler.py, then inspect the FSDP2 communication-fusion tests under test/qwen3_4b_fsdp2/. Run the existing unit and distributed tests described in the issue, including the HSDP gradient-scaling parity coverage. Done means the communication pipeline, memory handling, and numerical results match the stated Torch-aligned behavior without regressions.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python
- Domain
- distributed-systems, machine-learning, performance
- Issue type
- Feature
- Difficulty
- 5/5
- Estimated time
- Over a week
- Activity status
- Stale
- Clarity
- Needs clarification
- Newbie friendliness
- 25/100