大模型系统中的通算融合:从 overlap 到 tile-level operator
0. 范围、版本和证据
这里的通算融合不是一个单独的算子。实际代码里通常是几层东西叠在一起:上层改执行顺序,中层把 collective 和 GEMM 切成 chunk/tile,底层再决定由 SM、Copy Engine、NIC 还是交换机推进通信。
本文按下面六条线展开:
- 框架级异步流水:bucket、prefetch、独立 stream 和 event。
- micro-batch 与执行图重排:用另一批计算覆盖当前通信。
- 算子级、tile 级融合:边收边算、边算边发。
- GPU 发起通信:symmetric memory、one-sided communication、persistent kernel。
- SM-free 与异构资源协同:通信尽量交给 Copy Engine、NIC,或只占少量 SM。
- 网内计算:把 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 到底在哪里执行。

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。

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 六层结构

越往下,粒度越细,通常也越依赖具体 GPU、通信库和拓扑;出错时更难定位。
2. 框架级异步流水
2.1 怎么做
框架级 overlap 通常有五个部件:
- 切分:把参数、梯度或 activation 切成 bucket/chunk。
- 就绪检测:autograd hook 或执行图知道某个 bucket 已经完整产生。
- 异步发起:在通信 stream 上调用异步 AllReduce、ReduceScatter 或 AllGather。
- 延迟等待:消费者真正读取结果前才插入 event/wait。
- 双缓冲或预取窗口:当前 buffer 被计算消费时,下一个 buffer 在通信。
只调用 async_op=True 不够。等待点要放到第一个消费者之前;如果 collective 后面马上 wait(),中间没有可覆盖的计算。

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()
对应代码主要位于:
megatron/core/distributed/distributed_data_parallel.pymegatron/core/distributed/param_and_grad_buffer.pymegatron/core/distributed/distributed_data_parallel_config.py
执行关系可以简化成:
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 大小的取舍
- 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 转换可能成为新的关键路径。

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]:

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

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

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

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

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

这组图对应的是粗粒度执行边界:先完成 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 工作。

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 比例不够。

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 化后变成:

要把这件事做起来,至少需要四个组件:
- 可由所有 rank 访问的 communication/user buffer。
- tile 粒度的 ready flag 或 barrier。
- 改造后的 GEMM tile scheduler,优先调度已经 ready 的 tile。
- acquire/release 或 system-scope memory ordering,保证看到 flag 时也能看到数据。
反方向 GEMM -> ReduceScatter 则在 epilogue 中暴露输出 tile 的完成状态,通信消费者逐 tile reduce/send。

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。

4.2.2 Dense MLP layer0:AllGather + GEMM
基线过程:
torch.distributed.all_gather(X_shard)
torch.matmul(X_full, W)
FLUX 的设计是:
- 在一个 stream 上通过 cudaMemcpy Engine 推进不同 rank 输入分片的复制。
- 同时启动经过改造的 CUTLASS GEMM。
- 每个 GEMM threadblock 根据自己需要的 M/K tile 等待对应 barrier,而不是等待整个 AllGather。
- tile scheduler 对 M tile 做 swizzle,优先计算本地已经存在或更早到达的数据块。
- 后续远端分片一旦 ready,相关 threadblock 继续计算。
关键源码:
src/ag_gemm/sm90_all_gather_gemm_tile_scheduler.hppsrc/ag_gemm/sm90_all_gather_gemm_tma_warpspecialized_*.hppsrc/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:
- threadblock 完成一个输出 tile。
- 根据 M 维 shard 布局确定目标 rank。
- 把 partial output 写入目标 rank 对应的 scatter buffer/区域。
- barrier/后续 reduction kernel 确认各 rank contribution 到齐并完成求和。
关键源码:
src/gemm_rs/epilogue_reduce_scatter.hppsrc/gemm_rs/epilogue_nvshmem_reduce_scatter.hppsrc/gemm_rs/reduce_scatter_kernel.hppsrc/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 的处理:
- routing 在外部产生
splits_gpu/splits_cpu和scatter_index。 - 通信逐步获得其他 rank 的 token。
- 沿独立的 M 维重排 grouped GEMM tile。
- 某个 expert 的某段 token 到达后,相关 M tile 可以立即进入 GEMM。
- 不需要等所有 rank、所有 expert 的 token 都到齐。
Hopper 的入口类为 GemmGroupedV3AGScatter,旧架构对应 GemmGroupedV2AGScatterOp。接口和完整示例位于 docs/moe_usage.md。

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={...},
)
源码位于:
transformer_engine/pytorch/module/base.pytransformer_engine/common/comm_gemm_overlap/transformer_engine/common/comm_gemm_overlap/userbuffers/
当前配置支持按 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 方法
- ring_exchange:AG/RS 被切成 TP-size 相关分片,P2P send/recv 与 GEMM chunk 流水;默认倾向使用 Copy Engine。
- pipeline:把 ReduceScatter 与 GEMM 分片流水,通常需要保留一部分 SM 给通信;当前代码不允许把该方法用于 AllGather。
- bulk/external:利用整块 collective 或另一个独立 GEMM 的执行窗口覆盖通信,粒度更粗。
默认映射不是所有 Linear 都相同。例如当前源码把 qkv_fprop、fc1_fprop 等放在 ring-exchange 类,把 proj_fprop、fc2_fprop 放在 pipeline RS 类,wgrad/dgrad 又使用 bulk 或 external。TP overlap 需要按前向、dgrad、wgrad 的数据依赖分别设计。

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_sm和set_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__.pycsrc/apis/mega.hppcsrc/jit_kernels/impls/sm100_fp8_fp4_mega_moe.hppdeep_gemm/include/deep_gemm/impls/sm100_fp8_fp4_mega_moe.cuhdeep_gemm/include/deep_gemm/scheduler/mega_moe.cuhdeep_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 保存中间激活。
这解决了两个问题:
- L1 和 L2 若完全分阶段,中间激活池需要容纳所有 token。
- L2 过早消费尚未产生的 block,会产生等待甚至调度死锁。
从源码看,get_num_l1_warmup_waves()、live pool block 上界和两级 task barrier 同时服务于性能和正确性;它们决定了 L2 什么时候可以消费 L1 的输出,也关系到会不会死锁。

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。

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 却读到旧数据。

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 最佳实践

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_sms、num_qps。
dispatch 的工作是:
- 根据路由计算每个目标 expert/rank 的 token 数和偏移。
- 在 NVLink 和 RDMA 域发送 activation、expert index、weight 等元数据。
- 在接收端形成按 expert 排列的 buffer。
- 返回
EPHandle,保存 combine 所需的反向映射和计数。
6.2.4 combine
combine 接收 expert GEMM 输出,根据 EPHandle 把结果发回源 rank,并按 token/top-k 语义归约。训练反向中 dispatch 和 combine 互为对偶:dispatch backward 本质上是 combine,combine backward 本质上是 dispatch。

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。

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
建议流程:
- 单独测 GEMM 和通信的 shape-specific latency。
- 从很小的
S_comm开始递增。 - 同时记录 GEMM 变慢比例和暴露通信下降比例。
- 以端到端 layer/step 最短为准,而不是通信带宽最高为准。
- 分别为 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 数量受硬件资源约束。

7.3 NVLink SHARP / NVLS
NVLS 把 reduction/multicast 放入 NVSwitch 域。NCCL 当前算法包括 NVLS 和 NVLSTree,通过 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。实现逻辑是:
- 创建目标 communicator 前设置
NCCL_COLLNET_ENABLE=1。 - 对目标 group 执行 barrier,强制 PyTorch/NCCL 完成 lazy communicator 初始化。
- 随后清除环境变量,避免后续 group 也抢占 SHARP 资源。
- 默认目标为 DP group;
dp_replica需要多个 distributed optimizer instance。
代码注释还说明实际集群可能需要在 Slurm 中设置 #SBATCH_NETWORK=sharp。use_sharp=True 只能表示尝试为该 communicator 启用,不能证明交换机最终执行了 SHARP。要结合 NCCL debug/profiler、交换机 telemetry 或对照 benchmark 验证。
7.6 分层 collective 怎么排
大规模系统通常组合节点内与节点间两层 collective:

例如 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 来覆盖通信 |

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。
第二步:先做框架级重叠
- 打开 gradient RS/AR overlap。
- 打开 parameter AG prefetch。
- 检查 bucket 顺序和大小。
- 对 MoE/PP 做 micro-batch 交错。
- 确认端到端收益,而不是只看时间线。
第三步:处理资源争用
- 测量通信 kernel 的 SM 数、occupancy 和 L2/HBM 压力。
- 尝试 CE、NVLS、SHARP 或少 SM 通信。
- 扫
num_sm/SM margin。 - 将 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 故障测试。

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 没变快,通常是计算被通信拖慢或额外同步/拷贝抵消了收益。

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. 小结
把这些实现放在一起看,可以得到几个比较实际的结论:
- 框架级 bucket/prefetch 应该先做。它改动小,收益也最容易测出来。
- micro-batch/执行图重排 适合 PP 和 MoE 训练,前提是确实能找到另一段独立计算。
- FLUX/COMET、TE UserBuffers 处理的是 AG/RS/A2A 和 GEMM 之间的细粒度依赖,重点在 tile 到达后马上计算。
- DeepGEMM Mega MoE 把 dispatch 到 combine 放进一个 kernel,代价是更强的架构、shape 和运行时约束。
- NVSHMEM、NCCL Device API、PyTorch Symmetric Memory 提供了在 GPU 侧读写 peer buffer、发 signal 和推进通信的接口。
- DeepEP、Copy Engine、NCCL GIN 解决的是通信占用多少 SM,以及能不能交给 CE/NIC 推进。
- SHARP/NVLS 把 reduction 放到交换机或 NVSwitch 域,端点 GPU 只负责提交和消费结果。
落地时我会按下面三个问题往下排:
能否通过调度找到独立计算?
如果不能,能否把依赖切到 tile 粒度?
通信推进能否从计算 SM 下沉到 CE、NIC 或交换机?
最后还是看端到端 step latency、tokens/s、MFU、显存和正确性,不能只看单个通信带宽或时间线重叠面积。
我也在自己的推理引擎里做这几类通信—计算重叠的实现和验证。具体的实现细节、适用的 shape 以及效率数据,放到后续文章里再展开。本文先把公开项目中的实现方式和判断方法整理清楚,不把后续的个人实测结果混进来。