作者:WS、PDX、ZCX、PZL from DeepLink Group @ Shanghai AI Lab
概述
在张量并行场景中,AllGather 往往位于 MatMul 的关键路径上:传统实现通常需要先完成集合通信,再启动矩阵乘法,通信与计算之间存在明显的串行等待。本文以 triton_all_gather_matmul.py 为例,解析如何借助 PyTorch Symmetric Memory 建立 GPU 间可直接访问的对称缓冲区,并通过细粒度信号将远端分片的到达过程与 Triton MatMul 的分块计算组成流水线。该方案利用 CUDA P2P/NVLink 数据通路降低通信对计算资源的干扰,在给定测试配置下获得约 1.56 倍加速。文章还将讨论同步协议、内存可见性及实现层面的局限。
Triton系列往期课程:
为什么要融合 AllGather 与 MatMul
在大模型张量并行中,一个完整矩阵通常沿某个维度切分到多张 GPU。当前 rank 若要执行后续矩阵乘法,往往需要先通过 AllGather 收集其他 rank 的输入分片,再对拼接后的完整张量进行计算。
朴素执行流程可以概括为:
发起 AllGather → 等待全部分片就绪 → 启动 MatMul
这种实现边界清晰,但也引入了一段全局等待:即使本地分片或部分远端分片已经可用,MatMul 仍需等待整个 AllGather 结束。融合优化的核心,是将通信粒度从"完整张量"细化为"可独立消费的数据块",使计算能够尽早启动,并与后续数据传输形成流水。
yifuwang 的示例实现使用 PyTorch Symmetric Memory 与 Triton 完成了这一设计。与主要依赖通信内核推进数据交换的传统集合通信方案相比,对称内存允许 GPU 通过 CUDA P2P 映射直接访问其他 rank 的缓冲区。在具备 NVLink 或可用 PCIe P2P 通路的机器上,数据可以沿 GPU 间链路传输,从而减少额外的数据搬运与主机参与,并为细粒度的通信---计算编排提供基础。
**需要强调的是,对称内存并不会自动带来性能提升。**实际收益仍取决于 GPU 拓扑、P2P 带宽、矩阵形状、分块策略以及通信能否被计算充分掩盖。
从集合通信理解数据流
ReduceScatter 与 AllGather
理解 AllGather 之前,可以先看它在数据流上互补的 ReduceScatter。

ReduceScatter 先对各 rank 的对应数据块执行归约,再将不同结果块分发给相应 rank。以 rank 2 为例,它最终保留所有 rank 第 2 个数据块的归约结果;其他 rank 的处理方式相同。图中的变量 Y 应分别代入 rank 0、rank 1、rank 2 和 rank 3 来理解。
AllGather 则执行相反方向的数据收集:每个 rank 提供一份本地分片,并在操作结束后获得所有 rank 分片按约定维度组成的完整结果。

仍以 rank 2 为例,其绿色分片会被其他 rank 收集,并写入各自输出缓冲区中与 rank 2 对应的位置。最终,每个 rank 都持有相同的完整数据集合。
在数据划分与归约运算匹配的前提下,ReduceScatter 后接 AllGather,与 AllReduce 在语义上等价。这种分解也经常用于理解或实现带有计算重叠的集合通信算法。
融合后的流水线
AllGather 与 MatMul 融合后,MatMul 不再把完整 AllGather 结果视为唯一输入边界,而是按通信块逐步消费数据:
-
本地分片无需通信,可以直接进入计算;
-
远端分片通过 P2P 通路搬运或映射访问;
-
每个通信块就绪后,通过进度信号通知对应 MatMul 线程块;
-
MatMul 仅等待当前所需的数据块,而不等待整个 AllGather 完成。
这一执行模式与 PyTorch PR #139227所展示的异步通信---计算融合框架一致:通信生产数据块,计算按依赖关系消费数据块,两者通过事件或进度信号解耦。

在这个框架中,AllGather 沿 M 维将输入分成若干 chunk(下文也称"通信块"),每个 chunk 又对应一组 MatMul 输出 tile。调度器不再固定按整个 M 维扫描,而是以 chunk 为单位选择下一批 tile;在返回这些 tile 前,先检查该 chunk 的就绪信号。因为 A @ B 的不同 M 行块可以独立计算,某个 chunk 一旦可读,其对应的输出 tile 就能立即开始计算,无需等待其他 chunk。
# Pivot tile_id so that M tiles are processed in their ready order.
# This pivot preserves the prior swizzling.
pid_m = (pid_m + NUM_PID_M_PER_COMM_BLOCK * RANK) % num_pid_m
comm_block_id = pid_m // NUM_PID_M_PER_COMM_BLOCK
if comm_block_id // NUM_COMM_BLOCKS_PER_RANK == RANK:
# Read from the local a_shard
offs_am_src = (pid_m * BLOCK_SIZE_M) % COMM_BLOCK_SIZE_M
a_ptr = a_shard_desc_ptr
else:
# Wait for and read from a_shard copied from remote ranks
wait_signal((progress_ptr + comm_block_id).to(te.uint64), flat_tid)
offs_am_sc = pid_m * BLOCK_SIZE_M
a_ptr = a_desc_ptr
上方代码展示了单个 MatMul tile 选择输入的过程。首先对 pid_m 做与当前 RANK 相关的循环偏移,让每个 rank 优先处理自己持有的 M 行块;随后用 comm_block_id 将计算 tile 映射到通信块。如果该通信块属于当前 rank,内核直接从 a_shard 读取,并用块内偏移定位数据;如果属于其他 rank,内核先在 progress[comm_block_id] 上等待,确认对应分片已经复制到本地聚合缓冲区 a 后再读取。因此,本地路径没有通信等待,远端路径也只会阻塞在当前 tile 真正依赖的分片上。

从整体上看,通信侧和计算侧共享同一套 chunk 编号。通信侧把远端 chunk 逐块写入本地聚合缓冲区,每完成一块就发布对应的 chunk_signals[i];计算侧的异步输入调度器用 tiles_per_chunk_m 建立 chunk 与 M 维 tile 的对应关系,并在取出远端 tile 前检查信号。当 tile_idx_pivot_m 被设为当前 rank 的本地起点时,各 rank 会先消费本地数据,同时错开 tile 起点,避免集中访问同一个远端 rank。随着信号依次到达,远端 chunk 的搬运与已就绪 chunk 的 MatMul 交叠执行;当所有 chunk 都被消费后,结果与"先完整 AllGather,再执行 MatMul"相同,但中间的全局等待被拆成了逐块依赖。
对称内存如何打通 GPU 间访问
分配对称缓冲区
示例首先通过 Symmetric Memory 分配本地输入分片:
a_shard = symm_mem.empty(
m // world_size,
k,
dtype=torch.bfloat16,
device=device,
)
这里每个 rank 持有形状为 (M / world_size, K) 的 a_shard。所谓"对称",并不意味着不同 GPU 上保存相同内容,而是各 rank 以一致的形状、数据类型和布局参与内存注册,使运行时能够建立稳定的跨 rank 地址映射。
a_shard = symm_mem.empty(
m // world_size, k, dtype=torch.bfloat16, device=device
).normal_()
a = torch.randn((m, k), device="cuda", dtype=torch.bfloat16)
b = torch.randn((k, n), device="cuda", dtype=torch.bfloat16).T.contiguous()
c = torch.randn((m, n), device="cuda", dtype=torch.bfloat16)
Rendezvous 建立访问关系
分配张量后,需要执行 rendezvous,让同一进程组中的 rank 完成内存注册与句柄交换:
def rendezvous(
tensor: torch.Tensor,
group: Union[str, "ProcessGroup"],
) -> _SymmetricMemory:
"""为进程组中的对称张量建立跨 rank 访问关系。"""
enable_symm_mem_for_group(group_name)
return _SymmetricMemory.rendezvous(tensor, group_name)
rendezvous 位于初始化阶段,属于barrier细粒度原语。完成后,每个 rank 可以通过返回的 handle 获取本地或远端缓冲区视图。实际数据传输仍由 GPU 发起,CPU 无需参与每个数据块的搬运与同步。
if mm_only:
rank = 0
world_size = int(os.environ.get("WORLD_SIZE", "8"))
else:
symm_mem_hdl = symm_mem.rendezvous("a_shard", group=dist.group.WORLD)
assert symm_mem_hdl is not None, "a_shard must be allocated via SymmetricMemory"
rank = symm_mem_hdl.rank
world_size = symm_mem_hdl.world_size
这一机制成立还依赖几个前提:所有参与者必须使用一致的内存布局;GPU 间需要具备可用的 P2P 访问能力;设备、进程组和张量生命周期也必须保持一致。若拓扑不支持直接 P2P,实际性能与可用路径可能显著不同。
数据搬运与细粒度同步
使用进度数组描述数据就绪状态
为了让计算端判断远端数据块是否可读,实现中在 GPU 上维护一个 uint32 进度数组:
progress = torch.zeros(
world_size,
dtype=torch.uint32,
device="cuda",
)
在完整实现中,进度项通常与 src_rank 和 split_id 一一对应。每个元素相当于一个轻量级 mailbox:生产者完成某个通信块的数据准备后写入 1,消费者在 MatMul 内核中轮询对应位置。
backend_stream = symm_mem._get_backend_stream(priority=-1)
if mm_only:
progress = torch.ones(world_size, dtype=torch.uint32, device="cuda")
else:
progress = torch.zeros(world_size, dtype=torch.uint32, device="cuda")
symm_mem_hdl.barrier(0)
backend_stream.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(backend_stream):
all_gather_with_progress(a_out, a_shard, progress, SPLITS_PER_RANK)
AllGather 侧按 world_size 和分片数迭代,依次获取远端缓冲区并发布完成信号:

# 获取指定 rank、指定分片的远端缓冲区视图
src_buf = symm_mem_hdl.get_buffer(
src_rank,
chunks[0].shape,
inp.dtype,
chunks[0].numel() * split_id,
)
# 发布该分片已经就绪的信号
symm_mem_hdl.stream_write_value32(
progress,
offset=src_rank * splits_per_rank + split_id,
val=1,
)
# 在确有全局阶段依赖时执行 barrier
symm_mem_hdl.barrier()
get_buffer 返回目标 rank 对称内存的张量视图,最后一个参数用于定位分片偏移;stream_write_value32 则按当前 CUDA stream 的顺序写入进度值。正确性依赖一个关键的 happens-before 关系:数据写入必须先于完成信号对消费者可见,否则消费者可能观察到信号,却读到尚未完成的数据。
barrier() 提供 rank 间的阶段同步,但它的粒度比单块就绪信号更粗。融合实现应尽量让热路径依赖细粒度信号,仅在初始化、缓冲区复用或阶段切换等确需全局一致性的地方使用 barrier,否则容易重新引入全局等待。
MatMul 按需等待远端分片
数据准备与信号发布后,Triton MatMul 内核按照 program ID 计算当前输出 tile 所依赖的通信块。本地分片可以立即读取;远端分片则先等待对应进度项,再从聚合缓冲区取数。
kernel = matmul_kernel_tma_persistent[grid](
desc_a_shard,
desc_a,
desc_b,
desc_c,
progress,
M,
N,
K,
BLOCK_SIZE_M=configs["BLOCK_SIZE_M"],
BLOCK_SIZE_N=configs["BLOCK_SIZE_N"],
BLOCK_SIZE_K=configs["BLOCK_SIZE_K"],
GROUP_SIZE_M=configs["GROUP_SIZE_M"],
COMM_BLOCK_SIZE_M=COMM_BLOCK_SIZE_M,
RANK=rank,
WORLD_SIZE=world_size,
FP8_OUTPUT=dtype == torch.float8_e4m3fn,
NUM_SMS=NUM_SMS,
num_stages=configs["num_stages"],
num_warps=configs["num_warps"],
)
log_triton_kernel(kernel)
comm_block_id = pid_m // NUM_PID_M_PER_COMM_BLOCK
if comm_block_id // NUM_COMM_BLOCKS_PER_RANK == RANK:
# 当前 tile 依赖本地分片,无需等待通信
offs_am_src = (pid_m * BLOCK_SIZE_M) % COMM_BLOCK_SIZE_M
a_ptr = a_shard_desc_ptr
else:
# 当前 tile 依赖远端分片,等待对应通信块就绪
wait_signal(
(progress_ptr + comm_block_id).to(tl.uint64),
flat_tid,
)
offs_am_src = pid_m * BLOCK_SIZE_M
a_ptr = a_desc_ptr
这里,BLOCK_SIZE_M 决定单个 Triton program 负责的 M 维 tile 大小,NUM_PID_M_PER_COMM_BLOCK 则建立计算 tile 与通信块之间的映射。通信块过大,会推迟首个远端 tile 的启动时间;通信块过小,则会增加信号、调度和地址计算开销。因此,这两个粒度需要结合矩阵形状与链路特性调优。
wait_signal 的实现与语义
wait_signal 只让线程块中的一个线程执行全局内存轮询,避免所有线程同时访问同一信号地址。观察到目标值后,再通过 CTA 级 barrier 唤醒整个线程块:
@triton.jit
def wait_signal(addr, flat_tid):
if flat_tid == 0:
tl.inline_asm_elementwise(
"""
{
.reg .pred %p<1>;
wait_block:
ld.global.relaxed.gpu.u32 $0, [$1];
setp.eq.u32 %p0, $0, 1;
@!%p0 bra wait_block;
}
""",
"=r, l",
[addr],
dtype=tl.int32,
is_pure=False,
pack=1,
)
tl.inline_asm_elementwise(
"bar.sync 0;",
"=r",
[],
dtype=tl.int32,
is_pure=False,
pack=1,
)
这段 PTX 的执行过程可以拆成两步:
-
flat_tid == 0的线程循环执行ld.global,直到进度值等于1; -
bar.sync 0确保同一 CTA 内的其他线程不会越过同步点,随后共同读取已经就绪的数据块并执行矩阵乘累加。
is_pure=False 用于阻止编译器将轮询访问当作可消除或可随意重排的纯计算。与此同时,ld.global.relaxed.gpu 的内存序较弱,生产端的数据发布顺序和作用域必须与之正确配合。工程实现不能只关注"信号值是否变化",还需要验证信号之前的数据写入已经对消费 GPU 可见。
当后续通信块的准备速度能够追上当前 MatMul tile 的计算速度时,数据搬运便可以被计算掩盖。若链路带宽不足或计算量过小,线程块仍会在 wait_signal 中停留,融合收益也会随之下降。
更深入的源码分析可参考《FusedAllGatherMatMul Triton 工程实现》。

与显式远端写接口的差异
NVSHMEM 提供了面向远端内存的 put、get 以及整数原子或信号类 API,远端写入语义相对直观。

在 Symmetric Memory 中,get_buffer 暴露的是 peer buffer 视图,数据"拉取"或"推送"通常通过普通张量 copy 表达:
# 建立对称内存并获取相邻 rank 的缓冲区视图
hdl = symm_mem.rendezvous(t, dist.group.WORLD)
peer_buf = hdl.get_buffer(next_rank, t.shape, t.dtype)
# Pull:从远端视图复制到本地张量
t.fill_(rank)
hdl.barrier(channel=0)
pulled = torch.empty_like(t)
pulled.copy_(peer_buf)
hdl.barrier(channel=0)
assert pulled.eq(next_rank).all()
# Push:将本地张量复制到远端视图
hdl.barrier(channel=0)
to_push = torch.full_like(t, rank)
peer_buf.copy_(to_push)
hdl.barrier(channel=0)
assert t.eq(prev_rank).all()
这种张量化接口易于与 PyTorch 算子组合,但底层传输方向、同步作用域和内存序不如专用通信 API 显式。开发者需要明确以下问题:copy 由哪张 GPU 发起、运行在哪条 stream、何时对远端可见,以及缓冲区何时可以安全复用。对于多 stream 或双缓冲流水,通常还需要额外的事件或版本化进度值,避免上一轮的 1 被下一轮误认为新信号。
性能结果与适用边界
原实现使用以下命令在单机 8 GPU 环境中运行:
torchrun \
--nnodes 1 \
--nproc-per-node 8 \
--rdzv-backend c10d \
--rdzv-endpoint localhost:0 \
--no_python python3 triton_all_gather_matmul.py \
--M 16384 \
--N 6656 \
--K 16384 \
--BLOCK_SIZE_M 128 \
--BLOCK_SIZE_N 256 \
--BLOCK_SIZE_K 64
在该矩阵形状与硬件配置下,融合版 AllGather + MatMul 相比基线获得约 1.56 倍性能提升。

从 timeline 可以看到,Memcpy 与 MatMul 在时间轴上形成了较稳定的重叠,说明分块就绪信号确实将原本串行的通信与计算组织成了流水线。

不过,1.56 倍是特定环境下的实验结果,不应直接外推到所有模型和集群。实际评估时至少应同时报告 GPU 型号与互联拓扑、PyTorch 与 Triton 版本、数据类型、预热次数、统计口径以及基线实现。还应分别检查端到端延迟、有效带宽、SM 利用率和等待信号占用的周期,以判断收益究竟来自通信隐藏、内核调度减少,还是其他实现差异。
总结
AllGather 与 MatMul 的融合,本质上是一次生产者---消费者流水线重构:Symmetric Memory 提供跨 GPU 的统一缓冲区访问能力,通信侧按块生产远端分片,Triton MatMul 按依赖关系消费本地与远端数据,并通过 GPU 端进度信号避免全局同步。
这一方案的价值不仅在于减少一次独立算子调用,更在于缩短数据从"局部可用"到"参与计算"的路径。要稳定获得收益,需要同时处理好通信块与计算 tile 的映射、数据写入与信号发布的内存序、进度状态复用以及硬件拓扑适配。在这些条件满足时,计算可以有效掩盖相当一部分 AllGather 开销;反之,细粒度同步本身也可能成为新的瓶颈。