Triton第四课:基于对称内存融合 AllGather 与 MatMul,提速1.56倍

作者: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系列往期课程:

第一课:不写 CUDA,也能实现高性能 GPU Kernel

第二课:实现LayerNorm算子的六种方式

第三课:Grouped GEMM 算子优化思路与实战技巧

为什么要融合 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_ranksplit_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 开销;反之,细粒度同步本身也可能成为新的瓶颈。

参考文献

  1. NVIDIA NCCL User Guide:Collective Operations

  2. yifuwang:triton_all_gather_matmul.py

  3. PyTorch PR #139227

  4. FusedAllGatherMatMul Triton 工程实现

相关推荐
zhangfeng11331 小时前
AMD Instinct MI50(gfx906)上为 Qwen 系列模型优化并可用的 vLLM 相关仓库、Docker 镜像与实践指南。
人工智能·docker·ai编程·qwen·算子开发·vllm·mi50
其实防守也摸鱼1 小时前
信创是什么:一文读懂信息技术应用创新
人工智能·阿里云·云计算·github·copilot
咖啡星人k1 小时前
2026 AI Agent 长期记忆:为什么 AI 助手总“聊完就忘“?MonkeyCode 免费上手
人工智能·深度学习·机器学习·语言模型·自然语言处理
帅哥的AI自修课1 小时前
AI 换个会话就装失忆?Letta 焊死「记忆即灵魂」,从 MemGPT 的 OS 梦到 Core/Recall/Archival 三层记忆一篇打通
人工智能
机构师1 小时前
AI编程实战:效率与成本,AI 编程的 ROI 怎么算
人工智能·prompt·ai编程·deepseek
ManageEngineITSM1 小时前
什么是CMDB?配置管理数据库的定义、作用与建设方法一文讲清
大数据·数据库·人工智能·资产管理·变更管理
Delite8021 小时前
开口闪点检测智能化升级:工业油品安全检测的标准化解决方案
大数据·人工智能·安全
CallFay云起未来1 小时前
AI客服能不能减少人工回复?从重复咨询到人机协同的落地分析
java·大数据·人工智能·文心一言
Elastic 中国社区官方博客2 小时前
OpenTelemetry Java 扩展:无需分叉 agent 即可自定义追踪
java·大数据·运维·开发语言·数据库·人工智能·elasticsearch