Mojo 在 NVIDIA Blackwell 上实现矩阵乘法 85% SOTA 性能:multicast、2xSM MMA 与流水线优化全解析

Mojo 在 NVIDIA Blackwell 上实现矩阵乘法 85% SOTA 性能:multicast、2xSM MMA 与流水线优化全解析 Mojo 在 NVIDIA Blackwell 上实现矩阵乘法 85% SOTA 性能multicast、2xSM MMA 与流水线优化全解析【免费下载链接】mojoThe Modular Platform (includes MAX Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo本文基于 Mojo 官方设计文档《Matrix Multiplication on Blackwell: Part 3——The Optimizations Behind 85% of SOTA Performance》整理聚焦于 Modular 平台 Mojo 语言在 NVIDIA BlackwellSM100GPU 上把矩阵乘法从约 300 GFLOPS 提升至 SOTA 85% 的三代内核优化路径CTA 级多播multicast与 2xSM MMA、2SM 流水线pipelining warp specialization以及写回write-out双缓冲。读者将掌握__llvm_metadata设定 CTA cluster、TMA 多播掩码计算、tcgen05.mma.cta_group::2指令、循环缓冲circular buffer与内存屏障mbar协作等实战技术并能在仓库对应的迭代式内核源码中逐一对照验证。系列背景与本文定位这是 Mojo 官方 Blackwell 矩阵乘法系列博客的第三篇。前两篇分别对应 Part 1Blackwell 架构与朴素内核和 Part 2共享内存分块、TMA、tcgen05.mma、swizzling、TMA store。在前两篇结束时4 行朴素内核的性能仅为 cuBLAS 的 0.3%经过分块 张量核 swizzling TMA store 后达到约 293.6 TFLOPS8.7% cuBLAS但性能仍被全局内存吞吐所限制。本篇的目标非常明确利用 Blackwell 的两项高级特性CTA 多播与 2xSM MMA并结合流水线与 warp specialization将性能提升约 5 倍达到 SOTA 的 85%。整篇文章围绕三个依次递进的内核展开Kernel 5引入 CTA cluster、TMA multicast 与tcgen05.mma.cta_group::22xSM MMA性能达到 360.2 TFLOPS约 20% SOTAKernel 6基于 2SM MMA 实现共享内存循环缓冲流水线并引入 warp specialization性能跃升至 1429 TFLOPS81% SOTAKernel 7对写回路径做双缓冲double-buffering额外获得 64 TFLOPS最终达到 85% SOTA。该系列的完整推导过程在仓库中有对应的可运行/可基准测试的迭代式内核源码matmul_blackwell_iterative 目录其中1_naive_sm100.mojo到8_clc_tmem_ping.mojo与博客系列的 Kernel 1~8 一一对应BUILD.bazel声明了这些测试仅在有 B200 GPU 的机器上运行。Kernel 5multicast 与 2xSM MMA从 Hopper 代开始NVIDIA GPU 的流处理器SM可以被分组同一 SM 组CTA cluster内的协作线程数组CTA能够互相访问彼此的共享内存即分布式共享内存访问。基于这一能力Blackwell 上出现了两个高级优化手段TMA multicastingHopper 起支持多个 SM 协作把一个 tile 加载进共享内存2xSM MMA2 个 SM 的张量核协作完成一次大型 MMA 运算输入直接来自各 SM 的共享内存。声明 CTA cluster__llvm_metadata使用这些特性的第一步是在编译期把 cluster 的尺寸告诉编译器。在 Mojo 中通过__llvm_metadata装饰器完成__llvm_metadata(nvvm.cluster_dimcluster_shape) def blackwell_tma_pair_umma_kernel[ a_type, ...other parameters..., cluster_shape ]上述代码在编译期设置了内核的 CTA cluster 发射形状。同一 cluster 内的 CTA 可以互相访问彼此的共享内存。在仓库源码 5_2sm.mojo 中可以看到完整的装饰器写法cluster 形状通过cluster_shape: StaticTuple[Int32, 3] StaticTupleInt32, 3作为内核的编译期参数传入默认(1, 1, 1)表示不使用 clusterKernel 5 的测试与基准分别使用cluster_shape(2, 1, 1)见 5_2sm.mojo 的 benchmark_blackwell_matmul。CTA 内存多播CTA memory multicasting为了直观理解多播文档使用了一个简单例子A、B、C 均为 256x256 矩阵发射 4 个 CTA每个 CTA 计算一个 128x128 的 tile。若 4 个 CTA 各自独立地从全局内存加载自己所需的 A、B tile会存在大量冗余加载——这与系列博客开头提到的元素级冗余类似但这次冗余发生在tile 粒度。解决办法是把这 4 个 CTA 组成一个 2x2 cluster由于 SM 可以访问彼此的共享内存可以让每个 SM 只从全局内存加载 tile 的一半然后广播broadcast给邻居。例如同一行的两个 CTA 各自只加载 A 矩阵 tile 的一半行另一半由对端提供B 矩阵同样按列切分广播。这样每个 CTA 的全局内存加载量减半却仍能获得完整的 tile。该技术称为 multicasting且可扩展——例如 4x4 cluster 时每个 CTA 只需加载自己 tile 的四分之一。主机侧按 cluster 维度切分 TMA tile在 Mojo 代码中首先在主机侧声明 TMA 算子。由于每个 CTA 会切分 A、B 的共享内存 tile 并广播给对应 peer CTA因此 TMA tile 的尺寸需要按 cluster 维度缩小a_tma_op create_tma_tile[ (BM // cluster_shape[1], BK), swizzle_modea_swizzle ](ctx, a) b_tma_op create_tma_tile[ (BN // cluster_shape[0], BK), swizzle_modeb_swizzle ](ctx, b)在 5_2sm.mojo 的主机封装函数blackwell_kernel_5中可以看到等价实现A 的 TMA tile 行数为BM // cluster_shape[1]列维切分B 的 TMA tile 行数为BN // (cluster_shape[0] // cta_group)行维按 cluster 与 cta_group 双重切分。设备侧async_multicast_load与多播掩码在内核中多播加载使用async_multicast_loadif elect_one_thread: .... a_tma_op.async_multicast_load( a_smem_slice, tma_mbar[0], (UInt(i * BK), UInt(a_gmem_slice_coord)), a_multicast_mask, )该 API 与上一篇文章中的async_load类似但有三个关键差异多播掩码a_multicast_mask这是一个 16 位值描述参与本次加载的 CTA 索引。每个 bit 代表一个 CTA一个 cluster 最多 16 个 CTA。例如在上述 2x2 cluster 的例子中CTA 0 和 CTA 2 负责多播 A tile因此掩码都是0b101。Mojo 中的计算方式为var rank_m block_id_in_cluster.x # CLUSTER_M 和 CLUSTER_N 是 cluster 的维度 comptime for i in range(CLUSTER_N): a_multicast_mask | 1 (i * CLUSTER_M) a_multicast_mask rank_m仓库 5_2sm.mojo 中同时计算了 A、B 两个掩码A 掩码按rank_m位移B 掩码还需结合peer_cta_coord[0]peer CTA 在 cluster 中的行内序号与rank_n * CLUSTER_M位移。共享内存切片a_smem_slice从原始 tile 的布局张量中按基础指针偏移切出alias a_tma_load_size a_desc_layout.size() var rank_n block_id_in_cluster.y var a_smem_slice type_of(a_smem_tile)( a_smem_tile.ptr rank_n * a_tma_load_size )全局内存坐标相应更新# 每个切片的行数 alias a_tma_rows a_desc_layout.shape[0].value() a_gmem_slice_coord block_idx.x * BM Int(rank_n) * a_tma_rows2xSM MMAtcgen05.mma.cta_group::2多播与分布式共享内存虽然减少了从全局内存到共享内存的数据量但 tile 在分布式共享内存中仍是重复存储的。例如在前面的例子里CTA 0 和 CTA 1 在共享内存中各保留一份BN x BK的 B tile 副本——既然每个 CTA 都能访问对方的共享内存保存两份显然是浪费。Blackwell 的 2xSM MMA 指令tcgen05.mma.cta_group::2正是为解决此问题而设计CTA 0 与 CTA 1 组成一对pair各自只加载 B tile 的一半2xSM MMA 指令能同时看到共享内存中的两半数据协调两个 SM 上的张量核完成一次相当于两个单 SM MMA 之和的大型 MMA。对比左右两图两个 SM 计算的是相同的2*BM x MMA_N x BK工作量其中MMA_N是 MMA 的 N 维尺寸、产生相同结果但 2xSM 指令把 B tile 的共享内存占用减半。与多播加载 单 SM MMA相比2xSM MMA 从两个层面减少共享内存流量其一CTA 0 和 1 仍各自加载 B tile 的一半与 Figure 5 的多播加载一致但这里只是普通 TMA 传输省略了多播步骤其二从共享内存到张量内存TMEM的 B tile 流量也减半因为分布式共享内存中只保存一份副本。发起 2xSM MMA 的代码启动 2xSM MMA 的代码与单 SM MMA 非常相似区别在于使用elect_one_cta由于两个 SM 协作完成一次 MMA只需其中一个 CTA每对 CTA 中 ID 为偶数的那个即 leader CTA发起指令if elect_one_cta: # 等待数据到达共享内存 ... if elect_one_thread: comptime for j in range(num_k_mmas): var c_scale 0 if i 0 and j 0 else 1 alias idx IntTuple(0, MMA_K * j) alias a_offset a_smem_layout(idx) * sizeof[a_type]() alias b_offset b_smem_layout(idx) * sizeof[b_type]() mmacta_group mma_arrive_multicastcta_group其中mma在cta_group2时调用tcgen05.mma.cta_group::2指令到达arrive函数也从上一篇的mma_arrive换为mma_arrive_multicast它接收cta_group参数并在cta_group2时向 leader CTA 的内存屏障发信号。在仓库 5_2sm.mojo 中2xSM MMA 被封装进MmaOpSM100_SS并以cta_group2参数化elect_one_cta block_rank_in_cluster() % 2 0L204用于选出 leader CTA。Tensor memoryTMEM布局使用 2xSM MMA 时两个 CTA 平分 M 维BM MMA_M / 2每个 CTA 在张量内存中持有形状为BM x MMA_N的一半结果。指令支持MMA_M取 128 与 256 两种值当MMA_M256时每一半的 TMEM 布局与单 SM MMA 相同上半部分存储于 leader CTA下半部分存储于其配对 CTA最终写回全局内存的方式与上一篇的 Kernel 4 一致当MMA_M128时布局不同文档省略了细节生产环境较少使用可直接查阅源码。性能小结采用多播与 2xSM MMA 后内核达到360.2 TFLOPS约 SOTA 目标的 20%。从 profile 看即使使用了高级 MMA 指令性能仍受限于全局内存吞吐——计算仍在等待数据传输完成。下一个优化将移除这个障碍增大计算与内存传输的重叠。Kernel 62SM 流水线pipelining为了保持代码整洁Kernel 6 把将 tile 加载进共享内存的逻辑封装为load_AB()把发起 MMA封装为consume_AB()写回结果封装为store_C()。以此粒度观察前一个内核的循环是for i range(K/BK): # 发起异步 TMA 加载 load_AB() # 等待屏障直到数据到达 tma_mbar[0].wait(tma_phase) # 发起异步 MMA consume_AB() # 等待屏障直到计算完成 mma_mbar[0].wait(mma_phase) store_C() # 写出结果在任意时刻有一半的硬件处于空闲要么张量核在等数据到达要么 TMA 在等 MMA 结束以便共享内存缓冲区可复用。数据依赖使得一个 CTA 无法同时利用两类硬件单元。流水化 MMA 与 TMA循环缓冲克服空闲的经典算法是在共享内存中引入多个缓冲区通过流水线让 MMA 与 TMA 重叠。首先要核算共享内存预算。在 Blackwell GPU 上内核最多可访问227 KB 共享内存共享内存与 L1 缓存合计提供 228 KB其中 1 KB 必须留给 L1。当前配置下BF16、最大 2xSM MMA 指令形状 256x256x16的占用为一个 A tileBM x BK x 2B (MMA_M / 2) * 64 * 2B 16 KB一个 B tileBN x BK x 2B (MMA_N / 2) * 64 * 2B 16 KB一个 C tileBM x MMA_N x 2B 64 KB合计还不到容量的一半。因此引入5 个流水级stages组织成循环缓冲随后让 TMA 与 MMA并行地遍历这些缓冲区当 MMA 消费一个缓冲区时另一个 TMA 可以预取prefetch后续计算所需的数据到另一个缓冲区从而显著提高计算与通信的重叠度。仓库 6_2sm_pipelined.mojo 的blackwell_kernel_6给出了共享内存的自动预算逻辑以 B200 上total_smem_size_available 233472字节为上限扣除 C tile 后按每级 A B 32 字节屏障开销计算max_pipeline_stages作为num_pipeline_stages传入内核smem_size再按 stage 数整体放大。这是文档中5 级流水的工程化版本——流水级数由共享内存预算动态推导。Warp specializationwarp 专业化要实现上述流水模式还需要另一个重要概念——warp specialization。此前所有内核只使用 4 个 warp源于从 TMEM 加载数据的需要且 TMA 与 MMA 都由线程 0 发起。为了让 TMA 与 MMA并行发起并作用于不同的流水级需要为不同 warp 分配不同任务if WarpRole.is_main_load(): for i in range(num_iters): ... mma_mbar[stage].wait(phase) # 发起 TMA 加载并信号 tma_bar ... if WarpRole.is_mma(): for i in range(num_iters): ... tma_mbar[stage].wait(phase) # 发起 2xSM MMA 并信号 mma_mbar ... mma_arrive_multicastcta_group即一个 warp 专门负责发起 TMA另一个 warp 专门负责发起 MMA二者并发地遍历各 tile。warp 之间通过内存屏障memory barrier互相告知tile 已到达或tile 已被消费、底层缓冲区可写入新数据MMA warp 等待 TMA 屏障由 TMA warp 发信号以确认输入就绪TMA warp 等待 MMA 屏障以确认缓冲区可写。仓库 6_2sm_pipelined.mojo 定义了WarpRole结构体MainLoad Self(4)、Mma Self(5)、Epilogue Self(3)并提供is_main_load()、is_mma()、is_epilogue()静态方法内核启动 6 个 warpblock_dim(32 * 6)见 L1007其中 1 个 TMA warp、1 个 MMA warp、4 个输出epiloguewarp。生产者与消费者分别使用独立的PipelineStateproducer_phase初始为(0, 1, 0)consumer_phase初始为(0, 0, 0)追踪循环缓冲的 stage 与 phase。输出阶段compute_barrier写回 TMEM 仍需要 4 个 warp。warp specialization 后为输出专门分配 4 个 warp同时需要一个新的内存屏障在 MMA warp 与输出 warp 之间传递MMA 结果已就绪的信号。整体结构演变为if WarpRole.is_main_load(): for i in range(num_iters): ... mma_mbar[stage].wait(phase) # 发起 TMA 加载并信号 tma_bar ... if WarpRole.is_mma(): for i in range(num_iters): ... tma_mbar[stage].wait(phase) # 发起 2xSM MMA 并信号 mma_mbar ... mma_arrive_multicastcta_group # 信号输出 warp if elect_one_sync(): mma_arrive_multicastcta_group if WarpRole.is_epilogue(): compute_barrier[].wait() # 将结果存储到全局内存 ...MMA warp 额外向新屏障compute_barrier发信号输出 warp 等待它以确保结果就绪随后执行与之前相同的存储逻辑。仓库 6_2sm_pipelined.mojo 的kernel_6主循环完整实现了这一结构mma_complete_mask由self_mask | peer_mask计算1 block_rank_in_cluster()与1 (block_rank_in_cluster() 1)。性能小结基准测试显示该流水线策略将性能提升至1429 TFLOPS即 SOTA 的 81%。Kernel 7写回双缓冲double-buffering the write-out此前优化一直聚焦于把数据加载进 matmul其实把结果存储到输出同样可以优化。先回顾当前把 TMEM 结果写回全局内存的方式def store_c(): # 把整个张量内存搬进寄存器 registers tcgen05_ld parameters # 把寄存器全部搬进共享内存 for block_offset in range(BN/TMA_BN): for st_matrix_offset in range(TMA_BN//16): st_matrix[c_smem_tile, registers] # 把整个共享内存 tile 搬进全局内存 c_tma_op.async_store( c_tma_tile, ((block_idx.x * **MMA_N** thread_idx.x * TMA_BN), (block_idx.y * BM)), )存储 C 数据包含三个主要步骤TMEM → 寄存器 → 共享内存 → 全局内存。由于数据依赖这些传输原本串行执行但 TMA store 是异步的这种串行并非必要。复用此前的思路可以对 TMA store 与 MMA 结果搬运做流水化。具体做法在共享内存中声明两个输出 tile称为 double-buffer形状为BM x StageN在二者之间 ping-pong 切换。为简单起见先从stageN 32开始该值可调。双缓冲流水策略如下stmatrix把第一个输出 tile 写入共享内存后立即发起 TMA store 但不阻塞转而从 TMEM 加载数据、用另一个缓冲区写共享内存。这样 TMA store 就与下一个 tile 的 TMEM → 寄存器 → 共享内存 传输重叠起来。在 Mojo 中实现时把输出按MMA_N / StageN 8次迭代拆分每次处理 TMEM 中的stageN 32列TMEM 加载与stmatrix代码与之前内核基本一致只是操作更小的 tile# 每次处理 32 列 alias stageN 32 alias num_stages MMA_N // stageN comptime for stage in range(num_stages): # 加载 TMEM ... # 用 stmatrix 存入共享内存 ... if elect_one_thread: # 发起 TMA store ... c_tma_op.commit_group() __parameter # 保持一个 TMA store 在途 if stage num_stages - 1: c_tma_op.wait_group[1]() # 最后一级等待所有 TMA store 完成 else: c_tma_op.wait_group[0]()代码的关键在于如何同步 TMA storecommit_group把此前发出的 TMA store 提交为一组wait_group[N]等待直到只剩最后 N 组仍在途第 N1、N2…组已完成。上述代码中第一次迭代不等待可以直接使用第二个缓冲区从第二次到倒数第二次迭代等待只剩 1 组 TMA store 在途从而让 TMA store 与下一轮的 TMEM 加载、stmatrix重叠最后一轮则等待所有 TMA store 完成。仓库 7_double_buf_writeout.mojo 的store_C完整实现了这套逻辑stageN c_smem_layout.shape[1].value()默认 32、num_stages MMA_N // stageNC 的共享内存迭代器c_iter.next(stage % 2)实现双缓冲 ping-pongnamed_barrier用于保证共享内存写入完成、TMA 读共享内存完毕在 stage 1 至 num_stages-2 之间守卫前一缓冲区wait_group[1]()/wait_group[0]()的同步逻辑与文档一致。同时内核参数num_output_stages: Int 2、output_tile_shape: IndexList[2] Index(128, 32)L494-L495直接暴露了双缓冲配置。一个免费的额外收益流水化带来一个附带好处。回顾 Kernel 5 的共享内存占用A、B tile 的 5 份流水副本5 * (BM BN) * BK * 2B 160 KB输出 C tileBM * BN * 2B 64 KB少量用于内存屏障、TMEM 地址的空间输出占用约 40% 的共享内存。而 Kernel 7 的输出只占2 * BM * StageN * 2B 16 KB省下的 48 KB 可以用来加深流水线为 TMA 与 MMA 提供更多重叠空间。性能小结该优化额外带来64 TFLOPS的提升内核达到SOTA 的 85%。下一步通向最后 15%至此还剩 15% 的差距。尽管 Kernel 6 已通过重叠 TMA 加载与 MMA 隐藏了全局内存加载开销但写回全局内存的开销仍然突出尽管 Kernel 7 做了一定优化同时后续 CTA 调度之间还存在启动开销launch overhead系列下一篇将用**持久内核persistent kernel**配合 Blackwell 的另一项高级特性——cluster launch controlCLC——同时解决这两个问题最终拉满 SOTA 差距。仓库中的 8_clc_tmem_ping.mojo 正是这一演进方向的工程实现结合 TMEM ping-pong 与 CLC可作为继续研读的入口。仓库源码地图本篇文章涉及的核心可验证资源均位于当前仓库内主题源码/文档路径系列前情Kernel 1~4分块、TMA、swizzle、TMA storematmul-on-blackwell-part-2.mdKernel 5multicast 2xSM MMA 完整实现与基准5_2sm.mojoKernel 62SM 流水线 warp specialization6_2sm_pipelined.mojoKernel 7写回双缓冲7_double_buf_writeout.mojo下一步persistent kernel CLC8_clc_tmem_ping.mojo迭代式内核的 Bazel 测试定义要求 B200 GPUmatmul_blackwell_iterative/BUILD.bazel这些测试文件均可通过mojo_test规则在有 B200 GPU 的环境仓库 Bazel 配置中//:has_gpu与//:b200_gpu约束下运行并支持--benchmark参数复现文档中的 TFLOPS 数据。例如 Kernel 5 的基准入口benchmark_blackwell_matmul会输出 Average time 与 Performance (TFLOPS)5_2sm.mojoKernel 6、7 的基准入口结构与之类似可直接对照文档中的 360.2 / 1429 / 149385% SOTA三组数据。总结本篇围绕 Blackwell 矩阵乘法性能优化的三条主线展开数据加载侧用 CTA 多播消除 tile 级冗余加载、用 2xSM MMA 消除分布式共享内存中的重复副本计算调度侧用循环缓冲 warp specialization 让 TMA 与 MMA 真正并行、并引入 compute_barrier 衔接 MMA 与输出 warp写回侧用双缓冲让 TMA store 与下一轮结果搬运重叠。三者叠加把内核性能从约 300 GFLOPS 一路推到 SOTA 的 85%。这些技巧在 Mojo 中均有对应的高层 APIasync_multicast_load、mma_arrive_multicast、PipelineState、commit_group/wait_group等且在仓库迭代式内核源码中一一可查是编写生产级 Blackwell 张量内核的绝佳范本。【免费下载链接】mojoThe Modular Platform (includes MAX Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考