在 300B MoE 的 MindSpore PyNative 训练中观察到,VPP Stage 的参数梯度归约位于上游 PP 反向梯度发送的关键路径上。HSDP 组内负载不均产生的 AllReduce 长尾会沿下游到上游逐级传播,最终表现为计算流上数百毫秒的 EVENT_WAIT 和 PP bubble。
EVENT_WAIT
本问题仅涉及 MindSpore 后端;Torch 后端不在本 issue 范围内。
相关但不重复的已有工作:
本 issue 聚焦 PP/VPP scheduler 与 HSDP reduction 的跨子系统调度顺序:即使 HSDP 内部 collective 已异步下发,只要其 wait() 落在 PP BWD_SEND 之前,仍会把参数梯度归约长尾暴露到 PP 关键路径。
wait()
BWD_SEND
pipeline_parallel_overlap_p2p: True
pipeline_parallel_overlap_b_f: True
pipeline_parallel_p2p_transport: multi_stream
profiling/pynative/4k/pp4_vpp2_4k_0912.tar.gz
Stage 分配:
pipeline_parallel_interleave_num: 2 pipeline_parallel_layers_per_stage: - "0-5,24-29" - "6-11,30-35" - "12-17,36-41" - "18-23,42-44"
旧版 add_fsdp_reduce_grad() 在检测到某个 virtual Stage 的最后一个 microbatch backward 后,立即插入 FSDP_REDUCE_GRAD:
add_fsdp_reduce_grad()
FSDP_REDUCE_GRAD
BWD(micro=last, stage=S) FSDP_REDUCE_GRAD(stage=S) BWD_SEND / BATCH_SEND_RECV 后续本地 VPP Stage 计算
例如 rank 12288 的 Stage 7 尾部近似为:
BWD(micro15, Stage7) FSDP_REDUCE_GRAD(Stage7) BATCH_SEND_RECV: BWD_SEND(micro15, Stage7) -> rank8192 BWD_RECV(micro12, Stage3) <- rank0 BWD(micro12, Stage3)
rank 8192 的 Stage 6 同理:
BWD(micro15, Stage6) FSDP_REDUCE_GRAD(Stage6) BATCH_SEND_RECV: BWD_SEND(micro15, Stage6) -> rank4096 BWD_RECV(micro12, Stage2) <- rank12288 BWD(micro12, Stage2)
stage.launch_reduce_grad() 不只是异步发起 collective,还会推进 HSDP RS/AR 状态机并等待/应用先前 pending 的 reduction。MindSpore 的 CommHandle.wait() 会在当前 device stream 上插入 aclrtStreamWaitEvent。旧逻辑没有切换 stream,因此 wait 落到主计算流(本例为 Stream 47),阻塞后续 batchSendRecv 的提交。
stage.launch_reduce_grad()
CommHandle.wait()
aclrtStreamWaitEvent
batchSendRecv
以用户看到的 rank 4096 事件为起点:
EVENT_WAIT on Stream 47 start: 5390.003955 ms duration: 474.035280 ms
它对应 rank 4096 与 rank 8192 之间的 PP hcom_batchSendRecv__982_37_1,约 8 MiB BFP16。实际 P2P 匹配完成只需约 0.5 ms,绝大部分时间在等下游 rank 8192 提交匹配通信。
hcom_batchSendRecv__982_37_1
完整依赖链:
因果关系:
rank12288: BWD -> HSDP AR wait -> BWD_SEND to rank8192 | rank8192: wait recv -> BWD -> HSDP AR wait -> BWD_SEND to rank4096 | rank4096: wait recv (474 ms) <---------+
rank 12288 的长尾涉及至少两类 replicate group:32-rank group 与 128-rank group。profiling 中 collective 的 Wait Time Ratio 接近 1,说明主要是 HSDP 组内到达不均衡,而非 8 MiB PP 传输带宽问题。当前 profiling 只采集了一条 PP replica,无法从该压缩包确定具体哪个 HSDP peer 是 straggler。
参数梯度与输入梯度存在两条不同依赖:
参数梯度:BWD -> HSDP ReduceScatter/AllReduce -> optimizer 输入梯度:BWD -> BWD_SEND -> 上游 Stage backward
BWD_SEND 只依赖 BWD 已产生的输入梯度 dx,不依赖参数梯度 AllReduce 完成。当前 scheduler 将二者不必要地串行化为:
dx
BWD -> HSDP reduction wait -> BWD_SEND
因此 HSDP collective 的 straggler tail 被放大成跨多个 PP rank 的关键路径气泡。现有 overlap_p2p=True 和 multi_stream 无法解决,因为发送端在 reduction wait 结束前尚未提交 P2P。
overlap_p2p=True
multi_stream
建议为 MindSpore PP+VPP+HSDP 增加 scheduler-owned 的异步 Stage reduction:
BWD完成 -> 优先提交关联的 BWD_SEND / batchSendRecv -> 在独立 MindSpore device stream 上推进该 Stage 的 HSDP reduction -> 主计算流继续后续本地 VPP Stage 计算和 P2P -> optimizer/最终梯度使用前,通过 device event 统一 drain
实现约束:
Stream
Event
compute_ready
reduce_done
reshard_after_backward
FSDP_RESHARD
在该 profiling 点,如果 rank 12288/8192 的参数梯度归约不再阻塞 BWD_SEND,rank 8192 的反向梯度有机会在 rank 4096 到达 5390 ms 等待点之前准备完成。该处 474 ms 的 PP wait 理论上可压缩到接近实际 P2P 匹配耗时(亚毫秒级)。
AllReduce 本身不会消失,实际端到端收益取决于它与后续 VPP chunk 计算、P2P 的覆盖比例和 HCCL 资源竞争,需要重新 profiling 验证,不能直接把 474 ms 全部计为端到端收益。
overlap_p2p=False
背景
在 300B MoE 的 MindSpore PyNative 训练中观察到,VPP Stage 的参数梯度归约位于上游 PP 反向梯度发送的关键路径上。HSDP 组内负载不均产生的 AllReduce 长尾会沿下游到上游逐级传播,最终表现为计算流上数百毫秒的
EVENT_WAIT和 PP bubble。本问题仅涉及 MindSpore 后端;Torch 后端不在本 issue 范围内。
相关但不重复的已有工作:
本 issue 聚焦 PP/VPP scheduler 与 HSDP reduction 的跨子系统调度顺序:即使 HSDP 内部 collective 已异步下发,只要其
wait()落在 PPBWD_SEND之前,仍会把参数梯度归约长尾暴露到 PP 关键路径。复现场景
pipeline_parallel_overlap_p2p: Truepipeline_parallel_overlap_b_f: Truepipeline_parallel_p2p_transport: multi_streamprofiling/pynative/4k/pp4_vpp2_4k_0912.tar.gzStage 分配:
pipeline_parallel_interleave_num: 2 pipeline_parallel_layers_per_stage: - "0-5,24-29" - "6-11,30-35" - "12-17,36-41" - "18-23,42-44"当前调度逻辑
旧版
add_fsdp_reduce_grad()在检测到某个 virtual Stage 的最后一个 microbatch backward 后,立即插入FSDP_REDUCE_GRAD:例如 rank 12288 的 Stage 7 尾部近似为:
rank 8192 的 Stage 6 同理:
stage.launch_reduce_grad()不只是异步发起 collective,还会推进 HSDP RS/AR 状态机并等待/应用先前 pending 的 reduction。MindSpore 的CommHandle.wait()会在当前 device stream 上插入aclrtStreamWaitEvent。旧逻辑没有切换 stream,因此 wait 落到主计算流(本例为 Stream 47),阻塞后续batchSendRecv的提交。Profiling 证据
以用户看到的 rank 4096 事件为起点:
它对应 rank 4096 与 rank 8192 之间的 PP
hcom_batchSendRecv__982_37_1,约 8 MiB BFP16。实际 P2P 匹配完成只需约 0.5 ms,绝大部分时间在等下游 rank 8192 提交匹配通信。完整依赖链:
因果关系:
rank 12288 的长尾涉及至少两类 replicate group:32-rank group 与 128-rank group。profiling 中 collective 的 Wait Time Ratio 接近 1,说明主要是 HSDP 组内到达不均衡,而非 8 MiB PP 传输带宽问题。当前 profiling 只采集了一条 PP replica,无法从该压缩包确定具体哪个 HSDP peer 是 straggler。
根因
参数梯度与输入梯度存在两条不同依赖:
BWD_SEND只依赖 BWD 已产生的输入梯度dx,不依赖参数梯度 AllReduce 完成。当前 scheduler 将二者不必要地串行化为:因此 HSDP collective 的 straggler tail 被放大成跨多个 PP rank 的关键路径气泡。现有
overlap_p2p=True和multi_stream无法解决,因为发送端在 reduction wait 结束前尚未提交 P2P。建议优化方案
建议为 MindSpore PP+VPP+HSDP 增加 scheduler-owned 的异步 Stage reduction:
实现约束:
Stream与Event;compute_readyevent,reduction stream 等待后再调用stage.launch_reduce_grad();reduce_doneevent,主流只在最终 optimizer/下一次必须消费梯度前等待;BWD_SEND应优先于大块 AllReduce 提交,避免 collective 抢占通信资源后继续延迟 PP 关键路径;reshard_after_backward、显式FSDP_RESHARD与 reduction stream 的所有权,避免同一 Stage 重复 reshard 或跨 stream 竞态;预期收益
在该 profiling 点,如果 rank 12288/8192 的参数梯度归约不再阻塞
BWD_SEND,rank 8192 的反向梯度有机会在 rank 4096 到达 5390 ms 等待点之前准备完成。该处 474 ms 的 PP wait 理论上可压缩到接近实际 P2P 匹配耗时(亚毫秒级)。AllReduce 本身不会消失,实际端到端收益取决于它与后续 VPP chunk 计算、P2P 的覆盖比例和 HCCL 资源竞争,需要重新 profiling 验证,不能直接把 474 ms 全部计为端到端收益。
验收标准
风险点
CommHandle.wait()绑定调用时的 current stream,必须保证所有 reduction 状态机 wait 都在 reduction stream 上执行;reduce_done,不能被后续 Stage reshard 或 optimizer 提前复用;