大模型系统中的通算融合:从 overlap 到 tile-level operator

0. 范围、版本和证据

这里的通算融合不是一个单独的算子。实际代码里通常是几层东西叠在一起:上层改执行顺序,中层把 collective 和 GEMM 切成 chunk/tile,底层再决定由 SM、Copy Engine、NIC 还是交换机推进通信。

本文按下面六条线展开:

  1. 框架级异步流水:bucket、prefetch、独立 stream 和 event。
  2. micro-batch 与执行图重排:用另一批计算覆盖当前通信。
  3. 算子级、tile 级融合:边收边算、边算边发。
  4. GPU 发起通信:symmetric memory、one-sided communication、persistent kernel。
  5. SM-free 与异构资源协同:通信尽量交给 Copy Engine、NIC,或只占少量 SM。
  6. 网内计算:把 reduction/multicast 下沉到 NVSwitch 或 InfiniBand Switch。

通信规避、低精度通信、zero-copy 和拓扑感知并行不算严格意义上的融合,但实际接入时绕不开,本文也会讨论。

0.1 源码快照

项目 检查的版本 看的内容
ByteDance FLUX 19831ca2d820e3e782ed1d15d8b52d0898b78b26 AG+GEMM、GEMM+RS、MoE/COMET 细粒度融合
DeepEP 01dc3aaac82068020353dce2c302e38153c0bfaa EP dispatch/combine、NCCL GIN、少 SM 通信
DualPipe 030ce4325f4ebeb437da4ebc6d00a70469dd58ae 双向 PP 调度与 F/B overlap
Megatron-LM 3d3323fbdafe697b6d67904a0bf9ad0dbef601ae DP/FSDP/TP/EP overlap 和 SHARP 接入
Transformer Engine 80ea3133efa4c3a3679845b8ee46dfc06e2792f0 UserBuffers、TP GEMM/collective overlap
DeepGEMM 559d79fb6994a58b8a15b4b93bf13ccc16edf247 Mega MoE 单 mega-kernel 融合

DeepEP V2 的接口变化比较快。重新对照代码时,先在项目根目录执行 git rev-parse HEAD,再检查接口差异。

0.2 性能数字怎么读

  • 源码确认:代码、README 或官方文档里能找到对应能力。
  • 作者报告:论文或项目 README 公布的数据,不是本机实测。
  • 这里没有在目标 GPU 集群上重新编译和测量 FLUX、DeepEP、DeepGEMM、NVSHMEM 或 SHARP。
  • 不同 GPU、NVLink 域、NIC、NCCL/CUDA 版本和模型 shape 下,性能结论不能直接外推。

0.3 先把四个概念分开

这四类动作经常一起出现,但改动的位置不同:

概念 主要改变什么 典型例子 是否一定需要自定义 kernel
调度重叠 执行顺序和等待点 bucket overlap、micro-batch、prefetch
算子融合 算子边界和数据粒度 AG+GEMM、GEMM+RS、A2A+Grouped GEMM 通常需要
资源卸载 通信由谁推进 Copy Engine、NIC、GPU-initiated communication 不一定
网内计算 reduction/multicast 在哪里执行 NVLS、SHARP、CollNet 否,但依赖硬件和运行时

所以,时间线重叠只能说明两个工作被同时提交过;算子融合要看依赖边界有没有改;SM-free 要看通信是不是还在跑 CUDA CTA;网内计算要看 reduction 到底在哪里执行。

结构图:concept-boundaries

1. 先看依赖关系

1.1 从串行时间到重叠时间

先看没有重叠时的执行顺序:

compute_0 -> communication_0 -> compute_1 -> communication_1

其时间近似为:

T_serial = T_compute + T_comm

如果通信能被另一段计算完全盖住,理想时间接近:

T_overlap = max(T_compute, T_comm) + T_sync

实际要看的不是时间线面积,而是暴露出来的通信时间:

T_exposed_comm = T_step - T_step_without_comm_side_effect

profiler 上两个 kernel 的时间区间重叠,不等于通信被隐藏。通信 kernel 可能和 GEMM 争抢 SM、L2、HBM 或 PCIe/NVLink 注入带宽,结果是 GEMM 变慢。判断有没有收益,还是要看 step latency、token throughput 或 MFU。

结构图:serial-overlap

1.2 两类依赖决定了实现难度

第一类是独立依赖:

计算 layer L-1 的梯度
        与
同步已经完成的 layer L 梯度

两者没有数据依赖,可以放到不同 stream。这是框架级 overlap 最容易处理的一类。

第二类是生产者—消费者依赖:

AllGather(X) -> GEMM(X, W)
GEMM(X, W)   -> ReduceScatter(Y)

后者不能按完整张量直接并发。要把张量和 GEMM 拆成 tile,把整个 AllGather 完成后才能算改成tile i 到达后就算 tile i。FLUX、COMET、Transformer Engine UserBuffers 和 DeepGEMM Mega MoE 都在处理这个依赖。

1.3 六层结构

结构图:layering

越往下,粒度越细,通常也越依赖具体 GPU、通信库和拓扑;出错时更难定位。

2. 框架级异步流水

2.1 怎么做

框架级 overlap 通常有五个部件:

  1. 切分:把参数、梯度或 activation 切成 bucket/chunk。
  2. 就绪检测:autograd hook 或执行图知道某个 bucket 已经完整产生。
  3. 异步发起:在通信 stream 上调用异步 AllReduce、ReduceScatter 或 AllGather。
  4. 延迟等待:消费者真正读取结果前才插入 event/wait。
  5. 双缓冲或预取窗口:当前 buffer 被计算消费时,下一个 buffer 在通信。

只调用 async_op=True 不够。等待点要放到第一个消费者之前;如果 collective 后面马上 wait(),中间没有可覆盖的计算。

结构图:async-flow

2.2 数据并行:梯度 bucket 与反向重叠

以 Megatron-LM DDP 为例,调用链是:

module parameter backward hook
  -> DistributedDataParallel._make_backward_post_hook()
  -> BucketGroup.register_grad_ready()
  -> 当前 bucket 的全部梯度 ready
  -> BucketGroup.start_grad_sync()
  -> async AllReduce 或 ReduceScatter
  -> iteration 末尾 finish_grad_sync()

对应代码主要位于:

执行关系可以简化成:

BWD layer L     -> grad bucket k ready -> launch RS/AR(bucket k)
BWD layer L-1   ------------------------^ communication in flight
BWD layer L-2   ------------------------^ communication in flight

结构图:bucket-overlap

bucket 大小的取舍

  • bucket 太大:collective 带宽效率高,但发起太晚,可覆盖窗口变小。
  • bucket 太小:发起得早,但小消息 latency、kernel launch 和协议开销变大。
  • bucket 顺序不匹配反向产生顺序:即使某些参数梯度已经完成,也会被尚未完成的参数挡住。
  • 多个 bucket 复用同一 communicator 时,发起顺序必须在所有 rank 上一致,否则容易死锁或串行化。

Megatron 还会根据 DP/CP group 大小和参数规模推导 bucket size。调优时要一起看 bucket size、NCCL channel/CTA 配置和反向层耗时,不能只改一个参数。

2.3 FSDP/ZeRO:参数预取与梯度 ReduceScatter

FSDP 的每个逻辑单元通常经历:

AllGather 当前单元参数
  -> forward/backward compute
  -> ReduceScatter 梯度
  -> 释放非本 rank 参数

这里的做法是把下一单元的 AllGather 提前:

compute(unit i)        ----------------------
AllGather(unit i + 1)     -------------------

PyTorch FSDP 的 BackwardPrefetch.BACKWARD_PRE 会在当前梯度计算期间预取下一组参数,覆盖更充分,但峰值时可能同时持有当前参数、下一组参数和当前梯度;BACKWARD_POST 重叠较少但显存压力更低。PyTorch FSDP 文档

Megatron 的参数同步路径使用 forward pre-hook:

start_param_sync()
  -> 异步 AllGather 参数 bucket
module forward_pre_hook
  -> finish_param_sync()
  -> 只在模块真正使用参数前等待

这里有几个限制:

  • prefetch distance 越大,不一定越快,可能因为额外参数常驻导致显存上涨或 allocator 压力。
  • AllGather 和 ReduceScatter 若在同一 NCCL communicator/内部 stream 上排队,发起顺序会影响实际并发。
  • 参数低精度存储、FP8/FP4 post-AllGather 转换可能成为新的关键路径。

结构图:fsdp-prefetch

2.4 Tensor Parallel 的入口

2.4.1 先看 TP 和 CP+TP 的数据布局

先把后面会反复出现的几个维度固定下来:S 表示 sequence/token 维,H 表示 hidden size,I 表示 FFN 的中间扩展维。下面的图不是 FLUX kernel,而是普通 GEMM、TP 和 CP+TP 的计算布局。先看清楚 collective 在哪里发生,再看 FLUX 如何把这些边界切到 tile 粒度。

普通 GEMM 的输入是 [S, H],权重是 [H, I],输出是 [S, I]

FLUX 基线:普通 GEMM

TP GEMM 沿权重的输出维切分。以两个 TP rank 为例,每个 rank 保存 [H, I/2] 的权重,计算得到本地的 [S, I/2],后续算子是否需要 AllGather 取决于它要消费完整结果还是继续使用 shard:

FLUX 基线:TP GEMM

CP+TP GEMM 先处理 sequence shard。每个 rank 只有 [S/2, H] 的 activation,先通过 AllGather 恢复到 [S, H],然后再按 TP 方式计算本地 [S, I/2]

FLUX 基线:CP+TP GEMM

FFN 可以看成两个 GEMM 中间夹一个 GeLU。普通 FFN 的第一层把 [S, H] 扩展到 [S, I],第二层再从 [S, I] 投影回 [S, H]

FLUX 基线:普通 FFN

TP FFN 中,两个 rank 各自计算一半的中间维。第二个 GEMM 的局部结果都是 [S, H] 的 partial output,最后通过 AllReduce 合并:

FLUX 基线:TP FFN

CP+TP FFN 则把两种切分叠在一起:输入先以 [S/2, H] 分布在各 rank 上,AllGather 后进行 TP FFN,最后通过 Reduce-Scatter 把结果重新分回 sequence shard [S/2, H]。图中采用的是两张 GPU、TP=2CP=2 的示意布局:

FLUX 基线:CP+TP FFN

这组图对应的是粗粒度执行边界:先完成 AllGather 或 GEMM,再启动下一个阶段。FLUX 要做的事情,是把这些边界进一步拆成通信 tile、GEMM tile 和 ready signal,让某个 tile 到达后就可以被下游计算消费,不必等待整个张量完成。

Megatron 的主要开关包括:

tp_comm_overlap
tp_comm_overlap_ag
tp_comm_overlap_rs
tp_comm_overlap_rs_dgrad
tp_comm_bulk_wgrad
tp_comm_bulk_dgrad

配置位于 megatron/core/model_parallel_config.py。这些开关只表达哪些线性层和 collective 需要重叠。实际的 chunk/tile pipeline 通常由 Transformer Engine UserBuffers 完成,见第 4 节。

2.5 能覆盖什么

这条方式适合:

  • 大 batch 训练;
  • 反向层数多、有较长独立计算窗口;
  • FSDP/ZeRO 参数和梯度通信;
  • 不希望维护自定义 CUDA kernel 的系统。

如果通信结果马上要作为 GEMM 输入,同时又没有其他 micro-batch 或其他层可以计算,框架级 overlap 就不够了,需要改执行顺序或进入 tile 级融合。

3. micro-batch 与执行图重排

3.1 用另一批计算盖住通信

一个 micro-batch 内部可能有严格依赖:

dispatch -> expert GEMM -> combine

两个 micro-batch 可以交错执行:

MB0: dispatch ------- expert GEMM ------- combine
MB1:          dispatch ------- expert GEMM ------- combine

这样 MB0 的专家计算可以盖住 MB1 dispatch,MB1 的专家计算可以盖住 MB0 combine。代价是要多保存几份 activation、路由元数据和通信 buffer。

3.2 Pipeline Parallel:GPipe、1F1B 和 DualPipe

GPipe 先执行所有 forward,再执行所有 backward,bubble 大且 activation 常驻时间长。1F1B 在 warmup 后交替执行 forward/backward:

warmup: F0 F1 F2 ...
steady: Fk B0 Fk+1 B1 ...
cooldown: ... Bn

1F1B 可以减少 pipeline bubble 和 activation 生命周期,但 P2P Send/Recv 是否真的和计算重叠,还要看:

  • isend/irecv 是否提前提交;
  • 通信 buffer 是否可以安全复用;
  • NCCL P2P 是否占用计算所需 SM;
  • 当前 stage 是否有独立的 forward/backward 工作。

结构图:pipeline-schedule

3.3 DeepSeek DualPipe

DualPipe 让 pipeline 从两端推进。每个 rank 持有两个 model chunk,两个方向的 micro-batch 在中间交汇,减少有效 pipeline 深度。

开源实现里主要看这些函数:

  • DualPipe.__init__() 接收两个 module chunk。
  • rank_mapping 把 process-group rank 映射到双向 PP rank。
  • _forward_compute_chunk()_backward_compute_chunk() 分别推进两个方向的 chunk。
  • _forward_backward_compute_chunk() 可以调用模型自定义的 overlapped_forward_backward(),在一个模块接口内共同执行 forward 与 backward。
  • _recv_forward()_send_forward()_recv_backward()_send_backward() 把 P2P 操作累计到 comm_ops
  • _commit_and_wait_comm() 使用 dist.batch_isend_irecv() 批量提交并等待。

源码入口:

DualPipe 提供的是调度骨架。README 要求真实模型实现自己的 overlapped_forward_backward();没有这个接口时,代码会先 forward、再 backward。用了 DualPipe 的调度类,不代表 F/B kernel 已经在 GPU 上并发。

DualPipe 的主要代价:

  • 每设备保存两个模型 chunk,参数占用约为传统单 chunk PP 的 2 倍。
  • activation 生命周期和数量随 schedule 改变。
  • 要拆分 weight-gradient 计算并管理 WeightGradStore,才能利用 zero-bubble 思路。
  • schedule 正确不等于通信已隐藏,仍需结合 Nsight Systems 检查 P2P 暴露时间。

3.4 DeepSeek MoE 的双 micro-batch

DeepSeek 公布的 V3/R1 profiling 说明 prefill 和 decode 都使用两个 micro-batch,使 All-to-All 与计算重叠;decode 路径还强调 RDMA 消息发出后尽量释放 GPU SM。DeepSeek profile-data

执行关系可以写成:

attention(MB0)  <-> EP communication(MB1)
expert(MB0)     <-> attention/communication(MB1)

Megatron 中对应的高层开关是 overlap_moe_expert_parallel_comm。当前源码还提供:

  • delay_wgrad_compute:延后权重梯度 GEMM,为通信制造覆盖窗口。
  • overlap_dispatch_backward_with_experts_wgrad:让 expert wgrad 与 EP A2A 并发。
  • ep_overlap_early_attn_memory_release:调整 attention backward 顺序,换取更低峰值显存,但可能暴露更多通信。

这些开关同时影响吞吐、显存和通信暴露时间。只看 overlap 比例不够。

结构图:microbatch

3.5 调度重排的边界

  • decode batch 很小,没有第二份足够大的独立计算。
  • MoE expert 严重不均衡,某些 rank 先完成、某些 rank 长时间拖尾。
  • attention 和 FFN 的计算/通信比例不匹配,无法相互完全覆盖。
  • micro-batch 继续切小后,Grouped GEMM 的 M 太小,计算效率下降超过通信收益。

这时就要看 tile 级融合,或者把通信交给 GPU/NIC 的其他资源。

4. 算子级、Tile 级融合

4.1 一般做法

假设 Y = AllGather(X) @ W。完整张量依赖是:

AllGather(X[0:P]) 全部完成 -> GEMM

tile 化后变成:

结构图:tile-ready

要把这件事做起来,至少需要四个组件:

  1. 可由所有 rank 访问的 communication/user buffer。
  2. tile 粒度的 ready flag 或 barrier。
  3. 改造后的 GEMM tile scheduler,优先调度已经 ready 的 tile。
  4. acquire/release 或 system-scope memory ordering,保证看到 flag 时也能看到数据。

反方向 GEMM -> ReduceScatter 则在 epilogue 中暴露输出 tile 的完成状态,通信消费者逐 tile reduce/send。

结构图:tile-boundary

4.2 ByteDance FLUX

这里的 FLUX 指 ByteDance Seed 的通信—计算融合 kernel 库,不是文生图模型 FLUX。项目入口为 bytedance/flux,对应论文为

这里有两部分代码:

  • FLUX 论文/基础能力:重点解决 Tensor Parallel 下 dense GEMM 与 AllGather/ReduceScatter 的依赖型重叠。
  • COMET/MoE 能力:在同一代码库中扩展 AllGather+Scatter+GroupedGEMM、GroupedGEMM+Gather+TopKReduce+ReduceScatter。

主要依赖:

  • CUTLASS:生成和改造高性能 GEMM mainloop/epilogue/tile scheduler。
  • NCCL:普通 collective 和 bootstrap 能力。
  • NVSHMEM:MoE 远端访问与跨 rank symmetric buffer;README 明确说明 MoE kernel 构建需要 NVSHMEM。
  • PyTorch/PyBind:Python operator 封装和框架集成。

4.2.1 FLUX 源码地图

功能 主要源码路径
Python API python/flux/src/pybind/
Dense AG+GEMM src/ag_gemm/
Dense GEMM+RS src/gemm_rs/
MoE AG+scatter+grouped GEMM src/moe_ag_scatter/
MoE grouped GEMM+gather+RS src/moe_gather_rs/
通信 collective src/coll/
shared-memory/NVSHMEM 封装 include/flux/ths_op/flux_shm.h
kernel registry 与 tuning src/generator/src/*/tuning_config/
Triton 实现 python/flux_triton/kernels/

设计说明见 docs/design.md,调优入口见 docs/tuning_guide.md

结构图:flux-map

4.2.2 Dense MLP layer0:AllGather + GEMM

基线过程:

torch.distributed.all_gather(X_shard)
torch.matmul(X_full, W)

FLUX 的设计是:

  1. 在一个 stream 上通过 cudaMemcpy Engine 推进不同 rank 输入分片的复制。
  2. 同时启动经过改造的 CUTLASS GEMM。
  3. 每个 GEMM threadblock 根据自己需要的 M/K tile 等待对应 barrier,而不是等待整个 AllGather。
  4. tile scheduler 对 M tile 做 swizzle,优先计算本地已经存在或更早到达的数据块。
  5. 后续远端分片一旦 ready,相关 threadblock 继续计算。

关键源码:

  • src/ag_gemm/sm90_all_gather_gemm_tile_scheduler.hpp
  • src/ag_gemm/sm90_all_gather_gemm_tma_warpspecialized_*.hpp
  • src/ag_gemm/ths_op/all_gather_gemm_op.cc

这里的融合不要求通信指令和 MMA 指令跑在同一个 warp 内。关键是通信和 GEMM 由同一个 operator 管理,GEMM 的 tile scheduler 能看到数据是否到达。

4.2.3 Dense MLP layer1:GEMM + ReduceScatter

基线过程:

Y_partial = X @ W_shard
Y = reduce_scatter(Y_partial)

FLUX 在 Ampere 路径中把 ReduceScatter 的数据分发逻辑放入 GEMM epilogue:

  1. threadblock 完成一个输出 tile。
  2. 根据 M 维 shard 布局确定目标 rank。
  3. 把 partial output 写入目标 rank 对应的 scatter buffer/区域。
  4. barrier/后续 reduction kernel 确认各 rank contribution 到齐并完成求和。

关键源码:

  • src/gemm_rs/epilogue_reduce_scatter.hpp
  • src/gemm_rs/epilogue_nvshmem_reduce_scatter.hpp
  • src/gemm_rs/reduce_scatter_kernel.hpp
  • src/gemm_rs/sm90_gemm_tma_warpspecialized_*_reduce_scatter.hpp

这样做消除了完整 GEMM 输出落地后再启动独立 RS的边界,也能减少中间 tensor 的额外读写。

4.2.4 MoE layer0:AG + Scatter + Grouped GEMM

普通 MoE 第一层包含:

AllGather tokens
  -> 根据 routing index scatter
  -> 按 expert 形成变长分段
  -> Grouped GEMM

FLUX/COMET 的处理:

  1. routing 在外部产生 splits_gpu/splits_cpuscatter_index
  2. 通信逐步获得其他 rank 的 token。
  3. 沿独立的 M 维重排 grouped GEMM tile。
  4. 某个 expert 的某段 token 到达后,相关 M tile 可以立即进入 GEMM。
  5. 不需要等所有 rank、所有 expert 的 token 都到齐。

Hopper 的入口类为 GemmGroupedV3AGScatter,旧架构对应 GemmGroupedV2AGScatterOp。接口和完整示例位于 docs/moe_usage.md

结构图:flux-moe

4.2.5 MoE layer1:Grouped GEMM + Gather + TopKReduce + RS

第二层的普通执行是:

每个 expert 做 GEMM
  -> 按原 token 顺序 gather
  -> 对 top-k expert 输出加权/归约
  -> ReduceScatter

FLUX 沿 N 维重新调度计算,并使用横向融合

  • 一组 CTA 负责高效率 GEMM。
  • 一组专用 CTA 负责 gather、top-k reduce 和通信。
  • GEMM tile ready 后通过 barrier 暴露给通信 CTA。
  • 通信 CTA 不需要等待整个 expert GEMM 完成。

调优参数 GatherRSHParams(gather_rs_ctas, n_dim) 中,第一个值就是专用于 gather/RS 的 threadblock 数量。太少会让通信跟不上,太多会抢占 GEMM 资源。

4.2.6 Ampere 和 Hopper 的差别

FLUX 设计文档特别区分了两代架构:

  • Ampere 常有远多于 SM 数量的 threadblock。某个 block 因远端 I/O 停顿时,SM 可以切换到其他 block,因此把远端写放进 epilogue 仍可能有效。
  • Hopper 高性能 GEMM 广泛采用 persistent warp specialization:producer warp 通过 TMA 搬运,consumer warp 执行 MMA,每个 SM 上常驻较少 block。若在 epilogue 中直接插入长延迟远端 I/O,会破坏原本紧凑的异步流水。

所以 Hopper 上通常会这样安排:

  • 计算和通信在同一个大 kernel/统一调度域中;
  • 但使用不同 warp-group 或专用 CTA 分工;
  • 通过 barrier 和 tile queue 连接,而不是让 MMA consumer 直接执行长延迟远端通信。

所以 kernel fusion 不是把所有代码塞进一个 CUDA kernel。要保住 GEMM mainloop 的效率,同时把粗粒度依赖拆掉。

4.2.7 FLUX 的调优维度

FLUX 不会对任意 shape 自动达到最优。要调的参数主要有:

  • GEMM tile_shape
  • cluster_shape
  • Cooperative/PingPong 等 kernel schedule;
  • raster order;
  • mainloop stage;
  • AG/RS chunk 数;
  • communication CTA 数;
  • 是否使用 per-tile flag、1D ring、P2P read、cudaMemcpyAsync;
  • dtype、FP8 fast accumulation 和 scale layout;
  • TP/EP size、节点数和实际拓扑。

项目的 tuning 流程会 profile 候选 kernel,把最优 GemmMeta + RuntimeConfig + GemmHParams 注册到 tuning config。生产接入时至少要覆盖真实模型的 (M,N,K,dtype,TP,EP,topk,arch) 组合,否则可能回退、找不到注册项,或选到不合适的默认配置。

4.2.8 FLUX/COMET 性能数字如何解读

FLUX 论文报告 fused kernel 最多可以覆盖 96% 通信,并报告训练和推理加速;COMET 论文报告单 MoE layer 和端到端加速。FLUX 论文COMET 论文

项目 docs/performance.md 也说明 Torch baseline 主要用于功能对照,没有充分优化。复现时至少要比较:

  • 当前优化版 PyTorch/NCCL;
  • Transformer Engine UserBuffers;
  • DeepEP + 高效 Grouped GEMM;
  • FLUX/COMET;
  • 相同 dtype、shape、拓扑、warmup、CUDA Graph 和路由分布。

4.2.9 FLUX 适合什么,不适合什么

适合:

  • TP 域位于高速 NVLink/NVSwitch 或稳定 P2P 域;
  • AG/RS 紧邻大 GEMM,粗粒度 overlap 无法隐藏;
  • MoE shape 相对稳定,允许离线 tuning;
  • 愿意维护架构特化 kernel。

不一定适合:

  • 极小 M,GEMM 本身没有足够工作覆盖通信;
  • 动态 shape 极多,tuning/注册成本高;
  • 跨节点网络抖动大,tile 到达次序难预测;
  • 需要完全通用的 collective 语义和快速迭代。

4.3 Transformer Engine UserBuffers

UserBuffers 的目标是为 Transformer 层的 TP 通信预先分配可跨 rank 使用的 buffer,并把每个 Linear 的 GEMM 与 AG/RS 绑定。

入口:

transformer_engine.pytorch.initialize_ub(
    shape=[sequence_length * batch_size, hidden_size],
    tp_size=tp_size,
    quantization_modes=[...],
    ub_cfgs={...},
)

源码位于:

当前配置支持按 GEMM 名称设置:

method
is_reduce_scatter
num_sm
cga_size
set_sm_margin
num_splits
aggregate
atomic_gemm
use_ce
fp8_buf
comm_priority / gemm_priority

三种 overlap 方法

  1. ring_exchange:AG/RS 被切成 TP-size 相关分片,P2P send/recv 与 GEMM chunk 流水;默认倾向使用 Copy Engine。
  2. pipeline:把 ReduceScatter 与 GEMM 分片流水,通常需要保留一部分 SM 给通信;当前代码不允许把该方法用于 AllGather。
  3. bulk/external:利用整块 collective 或另一个独立 GEMM 的执行窗口覆盖通信,粒度更粗。

默认映射不是所有 Linear 都相同。例如当前源码把 qkv_fpropfc1_fprop 等放在 ring-exchange 类,把 proj_fpropfc2_fprop 放在 pipeline RS 类,wgrad/dgrad 又使用 bulk 或 external。TP overlap 需要按前向、dgrad、wgrad 的数据依赖分别设计。

结构图:userbuffers

atomic GEMM 的含义

atomic_gemm=True 尝试让单个 GEMM kernel 在设备侧循环处理多个 chunk,减少 host 多次 launch;当前源码将它标为 beta/非全场景验证,并限制到 FP8。AG 和 RS 配对层必须使用一致的原子 ring-exchange 配置,否则 chunk 输出顺序不能正确还原。

UserBuffers 的限制

  • 需要 CUDA multicast 支持;否则可设置 UB_SKIPMC=1 尝试 CUDA IPC 路径。
  • buffer shape、TP domain 和量化模式必须在初始化时确定。
  • 需要针对具体 GPU/节点拓扑配置。
  • num_smset_sm_margin 配错时,可能出现时间线有 overlap,但总时间更慢
  • CUDA Graph、FP8 buffer、跨节点 TP 和不同版本 TE 的组合约束较多,必须用当前版本文档和代码验证。

4.4 DeepGEMM Mega MoE

Mega MoE 把整个 EP MoE 数据通路放进一个 mega-kernel,不再拆成几个独立 operator。

当前 README 定义的融合范围是:

EP dispatch
  -> Linear 1 (FP8 x FP4 或 BF16)
  -> SwiGLU
  -> Linear 2
  -> EP combine

源码入口:

  • DeepGEMM README Mega MoE
  • deep_gemm/mega/__init__.py
  • csrc/apis/mega.hpp
  • csrc/jit_kernels/impls/sm100_fp8_fp4_mega_moe.hpp
  • deep_gemm/include/deep_gemm/impls/sm100_fp8_fp4_mega_moe.cuh
  • deep_gemm/include/deep_gemm/scheduler/mega_moe.cuh
  • deep_gemm/include/deep_gemm/layout/mega_moe.cuh

Python 到 kernel 的调用链

get_symm_buffer_for_mega_moe()
  -> PyTorch distributed symmetric memory allocation/rendezvous
  -> 切分 input、scale、routing、workspace、combine buffer

transform_weights_for_mega_moe()
  -> 把两层 expert weight 转成 kernel 需要的 FP4/scale/TMA layout

fp8_fp4_mega_moe() 或 bf16_mega_moe()
  -> C++ 参数和 shape 校验
  -> heuristic 选择 block/pipeline/ring 配置
  -> JIT 生成并编译 SM100 kernel
  -> 单 kernel launch

内部调度机制

MegaMoEScheduler 把任务显式分成:

Linear1
Linear2
SharedLinear1
SharedLinear2

调度器先执行一定数量的 L1 warmup wave,确保 L2 消费某个 M block 前,对应 L1 任务已经发布。之后 L1/L2 交错执行,并使用 ring pool 保存中间激活。

这解决了两个问题:

  1. L1 和 L2 若完全分阶段,中间激活池需要容纳所有 token。
  2. L2 过早消费尚未产生的 block,会产生等待甚至调度死锁。

从源码看,get_num_l1_warmup_waves()、live pool block 上界和两级 task barrier 同时服务于性能和正确性;它们决定了 L2 什么时候可以消费 L1 的输出,也关系到会不会死锁。

结构图:mega-moe

dispatch 与 GEMM 如何连接

dispatch warp 从各 rank 的 symmetric buffer 拉取 token 和 scale,写入本地 expert pool,并更新每 expert 的接收计数。scheduler 轮询计数达到预期值,生成对应 expert 的 M-block 任务。这样某个 expert 的输入到齐后即可进入 L1,而不要求整个 All-to-All 完成。

L1、SwiGLU、L2 如何连接

  • TMA producer 把 activation、weight 和 scale 搬入 shared memory。
  • UMMA/Tensor Core consumer 执行 FP8×FP4 或 BF16 GEMM。
  • L1 epilogue 在 kernel 内直接执行 SwiGLU,并把结果放入循环复用的中间 pool。
  • L2 task 等待对应 pool block ready,随后执行第二次 GEMM。

combine 如何连接

L2 epilogue 根据 dispatch 阶段保存的 source rank、token index、top-k slot 元数据,把结果直接写到目标 rank 的 combine buffer。完成跨 rank barrier 后,kernel 内再读取 top-k slot、做归约并写最终输出。

从算子边界看,Mega MoE 去掉了这些中间步骤:

  • dispatch 与 L1 之间的 kernel 边界;
  • L1、SwiGLU、L2 的中间 launch 和部分中间显存流量;
  • L2 与 combine 之间的边界;
  • 多次 host synchronization。

约束

  • 当前主路径针对 SM100;不能把 README 的能力直接外推到任意 Hopper/Ampere。
  • 需要多进程 symmetric memory,README 示例要求 PyTorch 较新版本。
  • hidden/intermediate hidden、token 和 scale layout 有严格对齐要求。
  • 权重必须预转换成特定 FP4/UE8M0/TMA layout。
  • 一个 mega-kernel 包含大量 barrier 和跨 rank 状态,异常恢复、超时诊断和局部重试比独立算子困难。
  • 极端 expert imbalance 会造成部分 SM 长时间等待特定 expert 的 token。

4.5 三种实现怎么区分

方案 融合边界 主要场景 可插拔性 架构绑定
Transformer Engine UserBuffers 单个 Transformer Linear 与 AG/RS Dense TP 中到高
FLUX Dense AG+GEMM、GEMM+RS Dense TP
FLUX/COMET MoE MoE 两个子层分别融合 TP/EP MoE
DeepGEMM Mega MoE dispatch 到 combine 的整个 MoE 高度固定的 EP MoE 很高,当前主路径 SM100

选型时不能只看融合边界大小。融合越大,中间开销越少,但 shape 动态性、框架集成、调试和错误隔离都会变差。

5. GPU 发起通信

5.1 为什么要让 GPU 发起通信

传统流程:

CPU launch compute kernel
CPU launch NCCL kernel
CPU/CUDA stream 建立顺序
GPU 执行

当通信粒度降到 tile 或几十微秒时,host launch、CPU proxy、stream 边界和 CPU-GPU 同步开始占据明显比例。GPU-initiated communication 让 CUDA kernel 内部直接:

  • 读取或写入 peer GPU memory;
  • 发起 put/get;
  • 更新远端 signal;
  • 向 NIC 提交 RDMA work request;
  • 完成 reduce/copy building block。

结构图:device-initiated

5.2 symmetric memory 的编程方式

symmetric memory 要求参与 rank 以相同顺序创建相同大小/布局的 allocation。初始化后,每个 rank 能得到:

  • 本 rank buffer 指针;
  • peer rank 对应 buffer 指针或可映射地址;
  • signal pad;
  • 某些平台上的 multicast/multimem 指针。

典型 kernel 协议:

producer:
  write remote_data[tile]
  release/system fence
  signal remote_ready[tile]

consumer:
  wait/acquire remote_ready[tile]
  read remote_data[tile]
  compute(tile)

必须保证先数据、后信号的可见性。只使用普通原子加一而没有正确 scope/semantic,可能看到 ready flag 却读到旧数据。

结构图:symmetric-memory

PyTorch 已提供 alpha 状态的 torch.distributed._symmetric_memory,支持 CUDA、NVSHMEM、NCCL backend,以及 peer pointer、multicast pointer 和 signal pad。

5.3 NVSHMEM

NVSHMEM 把多 GPU 内存抽象成 PGAS,提供:

  • device-side put/get
  • put_signal、wait、barrier;
  • warp/block 粒度通信;
  • stream-ordered API;
  • NVLink/P2P 路径;
  • 跨节点 IBGDA/GDAKI 路径。

IBGDA 的重点是把 InfiniBand 控制面和数据面都放到 GPU,避免 CPU proxy。NVSHMEM 使用说明

QP 与线程映射

GPU 发起网络通信并非天然更快。需要调节:

  • DCI/RC QP 数量;
  • QP 按 CTA、SM、warp 或 DCT 映射;
  • 每次提交的请求批量;
  • 消息聚合大小;
  • 多 NIC/rail 映射。

每线程发送极小 put 会产生大量 WQE,可能比 CPU proxy 更差。官方建议根据模式使用单线程/线程块协作的 put,并尽量把零散数据打包成较大消息。NVSHMEM Device API 最佳实践

结构图:nvshmem-qp

5.4 NCCL Device API

NCCL 2.28 开始提供 device-side communication API,当前主要模块包括:NCCL Device-Initiated Communication

  • LSA:对 NVLink 或可 P2P 的 PCIe 设备做 load/store accessible 通信。
  • Multimem:利用 Hopper 及以后部分数据中心 GPU 的 NVLink SHARP multicast/reduction。
  • GIN:GPU-Initiated Networking,2.28.7 起支持网络路径。
  • reduce/copy building blocks:为自定义 compute-fused kernel 提供更底层构件。

它与传统 NCCL collective 的区别是:应用可以把通信 building block 嵌入自己的 CUDA kernel,不必把 ncclAllGather() 当成不可拆分的黑盒。

当前 DeepEP V2 已从 NVSHMEM 主 backend 切换到更轻量的 NCCL GIN,并复用已有 NCCL communicator。从这个实现可以看到,device API 已经开始进入实际的大模型通信库,而不只是底层实验接口。

5.5 persistent kernel

persistent kernel 长时间占据一组 CTA/SM,在设备侧循环:

poll task/doorbell
  -> issue network operation
  -> wait or process completion
  -> run compute tile
  -> publish output signal

优点:

  • 减少 kernel launch 和 CPU 调度;
  • 可以保持 routing/QP/peer pointer 状态;
  • 更容易做 tile 级 producer-consumer 流水。

风险:

  • 占用固定 SM,影响其他 kernel 调度;
  • grid-wide 或 cross-rank barrier 容易死锁;
  • CUDA Graph、抢占和故障恢复复杂;
  • 一个 rank 退出后,其他 rank 可能永久轮询。

6. SM-free 与异构资源

6.1 为什么 overlap 可能更慢

设独立执行时间为:

compute = 100 us
comm    = 60 us

理想 overlap 是 100 us。但如果通信 kernel 抢占 SM/L2,使 GEMM 从 100 us 变成 150 us,即使 60 us 通信完全藏在里面,总时间仍是 150 us,比合理调度更差。

所以这里要一起决定:

  • 通信由谁推进;
  • 占多少 SM;
  • 是否使用 Copy Engine/NIC;
  • GEMM 要不要主动留下 SM margin;
  • 通信和计算的 stream priority。

6.2 DeepEP

DeepEP 是面向 MoE Expert Parallel 的通信库,主要提供 dispatch 和 combine,而不是完整 expert GEMM。它的设计目标是:高吞吐或低延迟 All-to-All、低精度通信、少量 SM 使用,以及显式 overlap 接口。DeepEP 项目

6.2.1 V2 结构

V2 主要改了这些地方:

  • high-throughput 和 low-latency 接口统一到 ElasticBuffer
  • 主 backend 从 NVSHMEM 切到 NCCL GIN;
  • 复用 NCCL communicator;
  • 支持更大的 scale-up/scale-out EP domain;
  • 根据 expert/top-k 分析计算 SM 和 QP 数,而不是必须逐 shape autotune;
  • README 声明 V3-like legacy training 场景可用 4–6 个 SM 达到与 V1 24 SM 相当或更好的性能,这是作者报告,不是本文实测。

6.2.2 调用链

deep_ep.ElasticBuffer
  -> deep_ep/buffers/elastic.py
  -> csrc/elastic/buffer.hpp
  -> csrc/kernels/elastic/dispatch.hpp 或 combine.hpp
  -> deep_ep/include/deep_ep/impls/dispatch.cuh / combine.cuh
  -> csrc/kernels/backend/nccl.cu
  -> NCCL GIN / NVLink / RDMA

legacy NVSHMEM/IBGDA 路径仍在 csrc/kernels/legacy/docs/legacy.md 中,但不能把 V1 和 V2 的性能、依赖或配置混为一谈。

6.2.3 dispatch

输入通常包含:

  • token activation;
  • topk_idx
  • topk_weights
  • expert 数和 token 上界;
  • 可选 FP8 data + scale;
  • num_smsnum_qps

dispatch 的工作是:

  1. 根据路由计算每个目标 expert/rank 的 token 数和偏移。
  2. 在 NVLink 和 RDMA 域发送 activation、expert index、weight 等元数据。
  3. 在接收端形成按 expert 排列的 buffer。
  4. 返回 EPHandle,保存 combine 所需的反向映射和计数。

6.2.4 combine

combine 接收 expert GEMM 输出,根据 EPHandle 把结果发回源 rank,并按 token/top-k 语义归约。训练反向中 dispatch 和 combine 互为对偶:dispatch backward 本质上是 combine,combine backward 本质上是 dispatch。

结构图:deepep-flow

6.2.5 EventOverlap

调用:

recv_x, ..., handle, event = buffer.dispatch(
    ...,
    num_sms=num_comm_sms,
    async_with_compute_stream=True,
)

# 执行与 recv_x 无依赖的计算
independent_compute()

event.current_stream_wait()
expert_gemm(recv_x)

EventOverlap 的作用是把通信何时结束当前 compute stream 何时必须等待分离。它不会自动创造独立计算;调用方仍需要正确识别可以放在 wait 前的工作。

6.2.6 handle caching

decode 中如果路由布局在连续迭代间可复用,可缓存 EPHandle,避免重新计算 layout 和 CPU 同步。这对小 batch decode 很重要,因为几十微秒级额外 host round trip 可能直接落在延迟关键路径。

6.2.7 少 SM 和0 SM要准确理解

DeepEP V2 README 声明:

  • EP 通信使用少量可分析计算的 SM;
  • RDMA Engram 和 PP 提供 0 SM 实验路径;
  • CP 可使用 Copy Engine 达到 0 SM;
  • 旧版 0-SM RDMA low-latency EP 已不再支持。

0 SM通常表示数据移动/网络推进不需要持续运行 CUDA 通信 kernel,不表示整个操作没有 GPU 控制、同步、buffer 管理或 NIC 资源消耗。

6.2.8 能解决什么,不能解决什么

  • 它解决 dispatch/combine,不负责把 expert GEMM 本身做快。
  • 如果 dispatch 很快但 Grouped GEMM 因负载不均或小 M 很慢,端到端仍不会快。
  • num_sms 太少可能无法打满 NVLink/RDMA;太多会侵占 expert GEMM。
  • FP8 dispatch 减少字节数,但 scale 布局、精度和对端 kernel 接口必须一致。
  • V2 要求较新的 PyTorch、NCCL 和 CUDA;生产接入应固定版本矩阵。

6.3 Copy Engine collective

Copy Engine 是 GPU 上独立于 SM 的 DMA 引擎。对于纯复制型 AllGather/AllToAll,如果 buffer 具备 symmetric/P2P 条件,可以由 CE 搬运,让 SM 保留给 GEMM。

PyTorch Symmetric Memory 当前文档给出的 CE collective 条件包括:

  • NCCL 2.28 或更高;
  • GPU 间 P2P 可访问;
  • NCCL zero-CTA policy;
  • tensor 使用 symmetric memory allocation 并完成 rendezvous;
  • collective 使用异步/internal stream。

这种方式特别适合通信主要是复制、不需要复杂 reduction的路径。对于 ReduceScatter,归约计算仍需由 NVLS/SHARP、SM 或其他硬件完成,不能把所有 collective 都简单视为 DMA copy。

结构图:resource-offload

6.4 SM margin 的调优方法

设总 SM 数为 S,给通信保留 S_comm,计算使用 S-S_comm。需要测量:

T_total(S_comm) = max(
    T_compute(S - S_comm),
    T_comm(S_comm)
) + T_unhidden_sync

建议流程:

  1. 单独测 GEMM 和通信的 shape-specific latency。
  2. 从很小的 S_comm 开始递增。
  3. 同时记录 GEMM 变慢比例和暴露通信下降比例。
  4. 以端到端 layer/step 最短为准,而不是通信带宽最高为准。
  5. 分别为 prefill、decode、训练 forward/backward 建立配置,因为最优点通常不同。

7. 网内计算

7.1 普通 AllReduce 在端点做什么

普通 ring/tree AllReduce 中,GPU/NIC 除了发送数据,还要在端点执行 reduction、转发中间结果并参与多轮同步。规模扩大后会出现:

  • GPU SM 被通信 kernel 占用;
  • 相同数据多次穿过端点;
  • 多轮网络 hop;
  • straggler 和网络 jitter 被 collective 放大。

7.2 InfiniBand SHARP

SHARP 把 AllReduce/Reduce/Broadcast 等集体操作的聚合逻辑下沉到交换机 ASIC:

GPU/NIC leaf contribution
   -> switch 中逐级 reduce
   -> 聚合结果向下分发

这样可以减少端点传输量、GPU reduction 工作和同步抖动。NVIDIA SHARP 介绍

SHARP 不是任意可编程的 GPU kernel,限制包括:

  • 支持的是有限 collective、dtype 和 reduction op;
  • 需要相应 InfiniBand switch/NIC/firmware/SHARP manager;
  • 集群调度系统可能需要显式申请 SHARP 资源;
  • 多租户和 process group 数量受硬件资源约束。

结构图:sharp-flow

7.3 NVLink SHARP / NVLS

NVLS 把 reduction/multicast 放入 NVSwitch 域。NCCL 当前算法包括 NVLSNVLSTree,通过 NCCL_NVLS_ENABLE 自动或显式启用。NCCL 环境变量文档

它适合:

  • 节点内或 rack-scale NVLink domain;
  • AllReduce、ReduceScatter、AllGather 等规则 collective;
  • 需要降低 NCCL SM 占用的 TP/DP 通信。

NCCL Device API 的 multimem 还允许自定义 kernel 使用 NVLink SHARP 提供的 multicast/reduction memory 语义,把网内能力进一步暴露给算子融合。

7.4 NCCL buffer registration 与 zero-copy

SHARP/NVLS 要发挥效果,buffer 注册很重要。注册后的 user buffer 可以直接参与 collective,减少 NCCL 内部中转拷贝和 channel/SM 资源使用。NCCL User Buffer Registration

所有 rank 必须一致使用注册或未注册 buffer,混用会产生未定义行为。CUDA Graph registration 和 local registration 的生命周期也需要与模型 buffer 复用策略一致。

7.5 Megatron-LM 如何接入 SHARP

当前 initialize_model_parallel() 支持:

use_sharp=True
sharp_enabled_group="dp" 或 "dp_replica"

源码位于 megatron/core/parallel_state.py。实现逻辑是:

  1. 创建目标 communicator 前设置 NCCL_COLLNET_ENABLE=1
  2. 对目标 group 执行 barrier,强制 PyTorch/NCCL 完成 lazy communicator 初始化。
  3. 随后清除环境变量,避免后续 group 也抢占 SHARP 资源。
  4. 默认目标为 DP group;dp_replica 需要多个 distributed optimizer instance。

代码注释还说明实际集群可能需要在 Slurm 中设置 #SBATCH_NETWORK=sharpuse_sharp=True 只能表示尝试为该 communicator 启用,不能证明交换机最终执行了 SHARP。要结合 NCCL debug/profiler、交换机 telemetry 或对照 benchmark 验证。

7.6 分层 collective 怎么排

大规模系统通常组合节点内与节点间两层 collective:

结构图:hierarchy

例如 ReduceScatter 可以先在 NVLink domain 内归约,再跨节点 SHARP,最后在节点内分发。分层方案的关键是 rank placement 必须匹配物理拓扑,否则逻辑分层反而产生跨 rail 或跨 NUMA 绕路。

8. 少做一些通信

8.1 并行策略和拓扑一起看

一般原则:

  • 高频、低延迟的 TP/CP 通信尽量限制在 NVLink/NVSwitch 域。
  • 大消息、频率较低的 DP 梯度同步可以跨节点。
  • EP placement 同时考虑 expert 热度、NVLink 域和 NIC rail。
  • PP 相邻 stage 尽量放在网络距离较近的 rank 上。

Megatron 源码也明确提醒相邻 rank 应尽量位于同一 DGX。逻辑 rank 编号不是无关细节,它决定 collective 和 P2P 的实际路径。

8.2 用 ReduceScatter + AllGather 代替 AllReduce

AllReduce 可以分解为:

ReduceScatter + AllGather

分解后的优势是:

  • ReduceScatter 可以和 backward 逐 bucket 重叠。
  • 每个 rank 只保留自己负责的梯度/参数 shard。
  • AllGather 可以在参数消费者前预取。
  • 更容易分别使用不同 communicator、SHARP 或 fused kernel。

但这不代表字节数凭空消失;收益来自 sharded state、调度窗口和更合适的 collective 边界。

8.3 Sequence Parallel

TP 下 LayerNorm、Dropout 等原本可能在各 rank 重复保存完整 sequence activation。Sequence Parallel 把序列/外层维度切分,减少 activation 常驻和部分重复计算,也把通信形式转成更容易与 GEMM 配对的 AG/RS。

8.4 低精度通信

典型做法:

  • gradient/parameter/activation 使用 FP8;
  • MoE dispatch 使用 FP8 data + per-token/per-block scale;
  • combine 根据精度要求使用 BF16、FP8 或其他压缩格式;
  • scale 与 payload 一同通信或预先共享。

收益是通信字节数下降,但必须计入:

  • quantize/dequantize kernel;
  • scale 计算和传输;
  • 对齐和 padding;
  • 数值误差与训练稳定性;
  • 下游 GEMM 是否能直接消费通信格式。

如果通信前转成 FP8、GEMM 前又还原成 BF16,量化和反量化会重新变成开销。更合适的做法是让接收 buffer 直接作为 FP8 GEMM 输入。

8.5 zero-copy 与布局融合

通信输出应尽量直接落到下游 operator 需要的布局:

  • expert-major/token-major;
  • TMA 对齐;
  • FP8/FP4 scale layout;
  • GEMM contiguous/grouped layout;
  • ReduceScatter 最终 shard layout。

如果通信之后还需要 index_select -> transpose -> pack -> copy,这些内存 kernel 往往重新成为关键路径。FLUX MoE 的 scatter/gather 融合、DeepEP 的新 GEMM layout、DeepGEMM Mega MoE 的预转换 weight 和 symmetric buffer 都在解决这个问题。

9. 怎么选

工作负载 第一优先级 第二优先级 不应首先做什么
大 batch DP 训练 gradient bucket overlap SHARP/分层 collective 直接写 mega-kernel
FSDP/ZeRO param prefetch + grad RS symmetric/user buffer 盲目增大 prefetch distance
节点内 Dense TP TE UserBuffers 或 FLUX NVLS/CE 只增加 NCCL stream
MoE 训练 DeepEP + 高效 Grouped GEMM + micro-batch overlap FLUX/COMET tile 融合 只优化 All-to-All 带宽
MoE prefill 高吞吐 dispatch/combine + 两 micro-batch COMET/Mega MoE 使用 decode 专用小消息配置
MoE decode 低延迟 GPU-initiated EP、handle cache、少 SM Mega MoE/细粒度 signaling 继续无限切 micro-batch
千卡跨节点 DP IB SHARP + topology-aware hierarchy bucket/schedule overlap 强行把 TP 跨慢网络扩大
极小 M、低延迟推理 persistent/device API、CE/NIC offload fused small-M operator 依赖大 GEMM 来覆盖通信

结构图:decision

10. 落地顺序

第一步:先做基线

至少记录:

  • GPU、SM 架构、NVLink/NVSwitch 拓扑;
  • NIC、rail、NUMA、PCIe 路径;
  • CUDA、driver、PyTorch、NCCL、TE/NVSHMEM 版本;
  • TP/PP/DP/EP/CP 配置;
  • 每层真实 (M,N,K,dtype)
  • collective 类型、消息大小、group size;
  • step latency、MFU、tokens/s 和 exposed communication。

第二步:先做框架级重叠

  1. 打开 gradient RS/AR overlap。
  2. 打开 parameter AG prefetch。
  3. 检查 bucket 顺序和大小。
  4. 对 MoE/PP 做 micro-batch 交错。
  5. 确认端到端收益,而不是只看时间线。

第三步:处理资源争用

  1. 测量通信 kernel 的 SM 数、occupancy 和 L2/HBM 压力。
  2. 尝试 CE、NVLS、SHARP 或少 SM 通信。
  3. num_sm/SM margin。
  4. 将 TP group 固定在高速域内。

第四步:给关键依赖做 tile 融合

先挑端到端占比高、shape 又比较稳定的边界:

AG + GEMM
GEMM + RS
MoE dispatch + GroupedGEMM
GroupedGEMM + combine

Dense TP 可以先评估 Transformer Engine UserBuffers;MoE 可以比较 DeepEP+GroupedGEMM、FLUX/COMET 和 DeepGEMM Mega MoE。

第五步:再考虑 mega-kernel

mega-kernel 放到最后。至少要满足下面这些条件:

  • shape、dtype、top-k、expert 布局稳定;
  • kernel launch 和中间内存流量占比明显;
  • 单算子融合仍留下同步边界;
  • 团队能维护架构特化 CUDA、barrier 和 symmetric memory;
  • 有可靠 correctness、超时、sanitizer 和多 rank 故障测试。

结构图:implementation

11. 怎么确认真的重叠了

11.1 Nsight Systems 看时间线

检查:

  • collective 是否提前发起;
  • wait 是否推迟到真实消费者;
  • compute 和 comm 是否实际并行;
  • GEMM 在 overlap 后是否显著变慢;
  • CPU 是否出现 launch gap;
  • 不同 rank 是否有长尾和不同步。

11.2 Nsight Compute 看资源

检查:

  • GEMM Tensor Core utilization;
  • SM occupancy 和 active warps;
  • L2/HBM throughput;
  • barrier stall、memory dependency stall;
  • 通信 CTA 占用的 SM 和寄存器/shared memory;
  • fused epilogue 是否破坏 mainloop pipeline。

11.3 做四组对照

至少做四组对照:

A. compute only
B. communication only
C. unfused serial compute + communication
D. fused/overlapped end-to-end

计算:

隐藏率 = 1 - (T_D - T_A) / T_B

同时报告 D 相对 C 的真实加速。若隐藏率很高但 D 没变快,通常是计算被通信拖慢或额外同步/拷贝抵消了收益。

结构图:validation

11.4 正确性测试

必须覆盖:

  • 不同 token 数和 expert imbalance;
  • 空 expert、极热 expert、top-k 边界;
  • 多 rank 不同到达顺序;
  • FP8/FP4 scale 和 padding;
  • CUDA Graph replay;
  • 多轮 buffer 复用和 sequence number 回绕;
  • 一个 rank 延迟或异常时的超时行为;
  • 与 BF16/reference 的误差阈值。

12. 容易混淆的地方

  • 误区一:两个 stream 就等于通算融合。 只有存在独立工作且硬件资源不严重冲突时,两个 stream 才能产生收益。

  • 误区二:通信 kernel 越快,端到端越快 通信为追求峰值带宽占用过多 SM,可能让 GEMM 变慢。DeepEP 的少 SM、TE 的 num_sm/set_sm_margin、CE collective 都是在优化端到端资源分配。

  • 误区三:把大 collective 切得越碎越好 切分会增加 signal、barrier、launch、QP/WQE 和协议头开销。最优 tile/chunk 需要结合网络 latency 和 GEMM tile 时间。

  • 误区四:FLUX、DeepEP、Mega MoE 是同一类库

  • FLUX/COMET 重点是把通信依赖融合进 GEMM/Grouped GEMM。

  • DeepEP 重点是高性能、少 SM 的 EP dispatch/combine。

  • Mega MoE 把完整 MoE 数据通路做成单个 mega-kernel。

​ 三者可以比较,也可能组合,但抽象边界不同。

  • 误区五:打开 SHARP 环境变量就证明用了网内计算 硬件、资源分配、communicator 创建顺序、NCCL 算法选择和 buffer 注册都可能导致回退。必须用 NCCL 日志、profiler 或网络 telemetry 验证。

13. 小结

把这些实现放在一起看,可以得到几个比较实际的结论:

  1. 框架级 bucket/prefetch 应该先做。它改动小,收益也最容易测出来。
  2. micro-batch/执行图重排 适合 PP 和 MoE 训练,前提是确实能找到另一段独立计算。
  3. FLUX/COMET、TE UserBuffers 处理的是 AG/RS/A2A 和 GEMM 之间的细粒度依赖,重点在 tile 到达后马上计算。
  4. DeepGEMM Mega MoE 把 dispatch 到 combine 放进一个 kernel,代价是更强的架构、shape 和运行时约束。
  5. NVSHMEM、NCCL Device API、PyTorch Symmetric Memory 提供了在 GPU 侧读写 peer buffer、发 signal 和推进通信的接口。
  6. DeepEP、Copy Engine、NCCL GIN 解决的是通信占用多少 SM,以及能不能交给 CE/NIC 推进。
  7. SHARP/NVLS 把 reduction 放到交换机或 NVSwitch 域,端点 GPU 只负责提交和消费结果。

落地时我会按下面三个问题往下排:

能否通过调度找到独立计算?
如果不能,能否把依赖切到 tile 粒度?
通信推进能否从计算 SM 下沉到 CE、NIC 或交换机?

最后还是看端到端 step latency、tokens/s、MFU、显存和正确性,不能只看单个通信带宽或时间线重叠面积。

我也在自己的推理引擎里做这几类通信—计算重叠的实现和验证。具体的实现细节、适用的 shape 以及效率数据,放到后续文章里再展开。本文先把公开项目中的实现方式和判断方法整理清楚,不把后续的个人实测结果混进来。

参考资料