尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

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

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, dtypetorch.bfloat16, devicedevice, )这里每个 rank 持有形状为(M / world_size, K)的a_shard。所谓“对称”并不意味着不同GPU上保存相同内容而是各 rank 以一致的形状、数据类型和布局参与内存注册使运行时能够建立稳定的跨 rank 地址映射。a_shard symm_mem.empty( m // world_size, k, dtypetorch.bfloat16, devicedevice ).normal_() a torch.randn((m, k), devicecuda, dtypetorch.bfloat16) b torch.randn((k, n), devicecuda, dtypetorch.bfloat16).T.contiguous() c torch.randn((m, n), devicecuda, dtypetorch.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, groupdist.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, dtypetorch.uint32, devicecuda, )在完整实现中进度项通常与src_rank和split_id一一对应。每个元素相当于一个轻量级 mailbox生产者完成某个通信块的数据准备后写入1消费者在 MatMul 内核中轮询对应位置。backend_stream symm_mem._get_backend_stream(priority-1) if mm_only: progress torch.ones(world_size, dtypetorch.uint32, devicecuda) else: progress torch.zeros(world_size, dtypetorch.uint32, devicecuda) 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, offsetsrc_rank * splits_per_rank split_id, val1, ) # 在确有全局阶段依赖时执行 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_Mconfigs[BLOCK_SIZE_M], BLOCK_SIZE_Nconfigs[BLOCK_SIZE_N], BLOCK_SIZE_Kconfigs[BLOCK_SIZE_K], GROUP_SIZE_Mconfigs[GROUP_SIZE_M], COMM_BLOCK_SIZE_MCOMM_BLOCK_SIZE_M, RANKrank, WORLD_SIZEworld_size, FP8_OUTPUTdtype torch.float8_e4m3fn, NUM_SMSNUM_SMS, num_stagesconfigs[num_stages], num_warpsconfigs[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 %p1; wait_block: ld.global.relaxed.gpu.u32 $0, [$1]; setp.eq.u32 %p0, $0, 1; !%p0 bra wait_block; } , r, l, [addr], dtypetl.int32, is_pureFalse, pack1, ) tl.inline_asm_elementwise( bar.sync 0;, r, [], dtypetl.int32, is_pureFalse, pack1, )这段 PTX 的执行过程可以拆成两步flat_tid 0的线程循环执行ld.global直到进度值等于1bar.sync 0确保同一 CTA 内的其他线程不会越过同步点随后共同读取已经就绪的数据块并执行矩阵乘累加。is_pureFalse用于阻止编译器将轮询访问当作可消除或可随意重排的纯计算。与此同时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(channel0) pulled torch.empty_like(t) pulled.copy_(peer_buf) hdl.barrier(channel0) assert pulled.eq(next_rank).all() # Push将本地张量复制到远端视图 hdl.barrier(channel0) to_push torch.full_like(t, rank) peer_buf.copy_(to_push) hdl.barrier(channel0) 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 开销反之细粒度同步本身也可能成为新的瓶颈。参考文献NVIDIA NCCL User GuideCollective Operationsyifuwangtriton_all_gather_matmul.pyPyTorch PR #139227FusedAllGatherMatMul Triton 工程实现
返回列表