在 Mojo 中优化 Blackwell 矩阵乘法(二):TMA、Tensor Core 与 Swizzling 实战指南
【免费下载链接】mojoThe Modular Platform (includes MAX & Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo
本篇文章是 Mojo/MAX 开源仓库中 "Matrix Multiplication on Blackwell" 设计文档系列的第二部分(关联文档见 matmul-on-blackwell-part-2.md),主题是从一个仅达到 cuBLAS 0.3% 性能的朴素 4 行 matmul kernel 出发,利用 NVIDIA Blackwell GPU 的硬件特性——共享内存、Tensor Memory Accelerator (TMA)、第五代 Tensor Core(tcgen05.mma)、Tensor Memory (TMEM) 与 Swizzling——逐步将性能提升 58 倍。读完本文,你将掌握基于 Mojo 的 GPU 内核开发中分块(tiling)、异步数据搬运、屏障同步、寄存器/共享内存数据布局等核心实战技能。
从 4 行 Kernel 出发:为什么性能只有 cuBLAS 的 0.3%
在第一篇文章(见 matmul-on-blackwell-part-1.md)中,我们介绍过 NVIDIA Blackwell GPU 架构,并最终得到一个 4 行的朴素 kernel。它的性能远逊于 cuBLAS——只有 cuBLAS 的0.3%,相当于把 1758 TFLOPS 的算力白白浪费掉了。本文将继续这一旅程,把性能提升到初始 kernel 的 50 倍以上。
为简化讨论,整个系列统一研究一个特定形状的矩阵乘法:A为MxK,B为KxN(已转置存储),结果C为MxN,且M=N=K=4096。先回顾朴素 kernel 的核心计算:
acc += a[row, k].cast[DType.float32]() * b[col, k].cast[DType.float32]()每次融合乘加(FMA)需要两次全局内存(GMEM)加载和一次内存写入。问题在于:全局内存虽然容量大,但远比其它层级的内存慢。因此,优化 matmul 的关键手艺就是借助 GPU 的内存层级结构,尽量规避或隐藏内存加载与存储。本文后续会用到的各类操作延迟对比如下:
再进一步,给每个线程分配一种颜色,可视化朴素 4 行 matmul 中各个线程对输入矩阵的读取方式:
- Thread 0 计算
C[0, 0],读取 A 的第 0 行与 B 的第 0 列 - Thread 1 计算
C[0, 1],读取 A 的第 0 行与 B 的第 1 列 - Thread 2 计算
C[1, 0],读取 A 的第 1 行与 B 的第 0 列 - Thread 3 计算
C[1, 1],读取 A 的第 1 行与 B 的第 1 列
仅看这 4 个线程就会发现:每个线程为了算一个输出值,要完整加载一行和一列;统计全部 4 个线程的内存加载次数,每一行和每一列都被重复加载了两次。这指向了我们能做的第一类优化:减少慢速全局内存的访问。
共享内存与循环分块(Loop Tiling)
减少冗余加载的经典技术叫循环分块(loop tiling)。思路很简单:把矩阵的一小块 tile 加载到快得多的缓存内存中,处理器在这一块数据上完成所有必要计算,不必频繁回到慢速主存;处理完一块后再加载下一块。
我们把共享内存(SMEM)当作这个缓存。每个 Blackwell SM 提供228KB 共享内存,因此多个线程可以在 block 内共享数据并做分块。
将矩阵划分为BMxBK的 tile(针对 A)和BNxBK的 tile(针对 B),其中BMxBNxBK = 64x64x64。关于可取数值及其限制后面再讨论,目前可以把这个大小视为4096x4096方阵的 tile。眼下唯一需要知道的约束是:tile 不能超过共享内存的大小。
在K/BK循环的第一次迭代中,从 A 加载一块BMxBKtile、从 B 加载一块BNxBKtile。这两块64x64tile 共8192个 2 字节元素,大约需要16KB共享内存,远小于 228KB 的可用容量。随后在这一 tile 上执行矩阵乘累加(MMA)运算,把结果存为中间值:
第二次迭代时,把接下来的两块数据载入共享内存,并累加本次 MMA 结果与上一次结果。如此循环K/BK次(本例为 256 次),直到得到最后一个 tile 的结果。K/BK循环结束后,输出 tile 就齐了,最终结果只向全局内存写一次:
Kernel 2:TMA 与 Tensor Core
第二个 kernel 比初始 kernel 更高级,同时使用分块和 Tensor Core 进行优化。粗略的骨架如下:
kernel_setup() for i in range(K // BK): load_tiles_ab() # leader thread loads A and B tiles issue_mma_axb() # leader thread issues MMA(A x B) transfer_c_tile_to_registers() # move final C tile from tmem to registers write_c_tile_to_global_memory() # store C tile from registers to gmemB 矩阵以转置形式存储,以保证访问时内存合并(coalesced)。这一步可以通过 Layout 变换完成:
alias a_layout = Layout.row_major(M, K) alias b_layout = Layout.row_major(N, K) # Transposed ...该 kernel 还需要一些 host 侧的设置改动,下文将逐步说明。
将 tile 加载进共享内存
NVIDIA Hopper 架构引入了Tensor Memory Accelerator(TMA)——一个专门的硬件单元,负责在全局内存(GMEM)与共享内存(SMEM)之间异步搬运数据。
使用 TMA 前,需要先在 host 侧创建一个 tensor tile(tensor map)并传入 kernel。tensor map 是一块128B 的数据块,编码了输入张量的形状、stride 和全局内存地址(它还可以编码swizzling 模式,稍后讨论)。在 Mojo 中用现成 API 创建 TMA tile:
# Rank 2 matrix # A/B tiles in shared memory have shapes BMxBK and BNxBK, respectively a_tma_op = create_tma_tileIndex(BM, BK) b_tma_op = create_tma_tileIndex(BN, BK)kernel 内使用 TMA 对象的方式如下:
alias num_iters = K // BK for i in range(num_iters): # One a single thread launches the TMA async copy. if elect_one_thread: tma_mbar[0].expect_bytes(expected_bytes) a_tma_op.async_copy( a_smem_tile, # shared memory tile containing the address tma_mbar[0], # barrier to guard the copy is finished (i * BK, block_idx.y * BM), # tile's coordinate in the input. ) b_tma_op.async_copy( b_smem_tile, tma_mbar[0], (i * BK, block_idx.x * BN), ) # All threads wait for the copy to finish. tma_mbar[0].wait(tma_phase) tma_phase ^= 1整体上,由单个线程(elect_one_thread)发起异步拷贝,并用内存屏障(tma_mbar)守护拷贝完成。a_tma_op.async_copy接受三个参数:
a_smem_tile:一个LayoutTensor,提供 tile 的共享内存地址;tma_mbar:用于跟踪已搬运数据量的内存屏障;(i * BK, block_idx.y * BM):当前 tile 在全局内存中的坐标,取决于迭代次数与 block 坐标。
为什么需要 TMA 屏障
由于 TMA 是异步操作,必须保证在 tile 完全落进共享内存之前 MMA 不能开始,否则数据竞争。这正是内存屏障(mbar)的用途:线程在屏障上等待/阻塞,直到 tile 复制完毕。
具体做法是给每个线程初始化自己的屏障相位(tma_phase=0),屏障内部也有自己的相位值(初始也为 0)。当线程的相位与屏障相位一致时,线程无法解锁屏障、无法通过;只有两者相位不同时线程才能继续:
tma_mbar[0].wait(tma_phase)执行时线程阻塞的过程如下:
在 TMA 传输开始前,用tma_mbar[0].expect_bytes(expected_bytes)告诉屏障预期接收多少字节。期望字节数即两块 tile 的字节总和:
alias a_expected_bytes = a_size * sizeof[a_type]() alias b_expected_bytes = b_size * sizeof[b_type]() alias expected_bytes = a_expected_bytes + b_expected_bytesTMA 会持续更新屏障已传输的字节数;一旦达到总量,屏障相位翻转,线程得以继续:
随后我们手动通过tma_phase ^= 1翻转每个线程的相位,保证线程在下一轮迭代中阻塞,直到那一次迭代的 tile 真正写入共享内存:
Mojo 也提供了抽象来隐藏部分 TMA 细节与优化技巧。例如,问:a_tma_op的 layout 是什么?答案是BMxBKtile,若BM=64、BK=64,其((shape), (stride))元组为((64, 64):(64, 1))(行主序/K 主序)。那么 TMA 单元加载这块64x64tile 需要多少次 fetch?答案是8 次,而非直觉上的 1 次——虽然我们指定了64x64的逻辑 tile 大小,TMA 硬件会把64x64分成 8 个64x8的子 tile 逐个加载。要解释原因,需要引入"core matrix"(核心矩阵)。
Core Matrices:TMA 的隐藏分块
TMA、Tensor Core 乃至整个 NVIDIA GPU 都存在一个隐藏细节:core matrix。
概念很简单:Tensor Core 不理解"元素",只理解"矩阵"。它只能把矩阵看作一组8x16B的 tile——即对我们来说8x8个元素的核心矩阵。
tcgen05.mma支持共享内存中 8 种规范化的 core matrix 布局(继承自 WGMMA),取决于布局(行主序或列主序)与 swizzle 模式。当前 kernel 对 A、B 均采用 K 主序,对应"每个 core matrix 的列(8x1)在共享内存中必须连续"的布局。这就是为什么描述符布局显示为(64, 8)——TMA 一次复制一列(8 个 core matrix),重复 8 次才能填满 tile 的 64 元素宽度:
当然,Mojo 库的async_copy把这些复杂度都抽象掉了:程序员只需发起一次拷贝,就可以期待 tile 出现在共享内存中。
发布 MMA 指令
回顾一下:Blackwell 引入的第五代 Tensor Core 带有一组新指令(tcgen05指令),对 MMA 操作有三项根本性改进:
- 单个 SM 上最大的
tcgen05.mma形状从 Hopper 的64x256x16提升到128x256x16,吞吐量几乎翻倍; - 引入 2SM
tcgen05.mma,最大可达256x256x16(2SM 操作将在本系列后续文章中解释); - 通过引入名为Tensor Memory的新型内存降低寄存器压力,
tcgen05.mma可以把结果存进 Tensor Memory 而非寄存器。那么 Tensor Memory 是什么?
什么是 Tensor Memory(TMEM)
TMEM 是一块256KB 的片上内存,专门用于存放tcgen05MMA 指令的输入或输出。它有 128 个 lane、每 lane 512 列,共 65,536 个元素;每个元素 4 字节,合计 256KB。分配按列进行,分配粒度为32 列,即一次最小分配 32 列(16KB):
在更早的 NVIDIA 世代中,矩阵乘结果必须存放在通用寄存器中,这带来几个问题:
- 寄存器空间稀缺(每个 SM 只有 64K 个寄存器),Tensor Core 与通用 ALU 之间存在争用;
- 寄存器是线程私有的,而前 Blackwell GPU 上 MMA 是 warp 级操作,因此发起 MMA 的 warp 必须等待其完成,才能继续依赖 MMA 结果的任务(如 epilogue)。
TMEM 解决了这些问题:把 ALU 使用的寄存器与 Tensor Core 所需的寄存器彻底分离。
在代码中这样使用tcgen05.mma和 tensor memory:
for i in range(num_iters): load_tiles_ab() #section 1 if elect_one_thread: comptime for j in range(num_k_mmas): 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]() # Use c_scale=0 for the first mma to initialize results and use # c_scale=1 subsequently to accumulate results. var c_scale_value: UInt32 = 0 if (i == 0 and j == 0) else 1 mma( adesc + a_offset, bdesc + b_offset, tmem_addr, idesc, c_scale=c_scale_value, ) mma_arrive(mma_mbar) mma_mbar[0].wait(mma_phase) mma_phase ^= 1tcgen05.mma指令异步执行,与 TMA 操作类似——由单一线程发起、由内存屏障守护。区别在于这里用mma_arrive(包装了tcgen05.commit)来发信号给内存屏障,并把它与正在执行的 MMA 指令动态关联。
注意我们发布了num_k_mmas条 MMA 指令(而不是把 A、B 两个 tile 一次性喂给 Tensor Core 相乘)。原因是BMxBNxBK的分块并不够——真实硬件指令有尺寸限制:tcgen05.mma要求 K 维度为 32B(即 BF16/FP16 的 16 个元素)。因此BK=64时 MMA 需要 4 次迭代:
没错,这实际上是一个嵌套分块策略。mma函数调用tcgen05.mma指令,结果累加到地址tmem_addr指向的 tensor memory 中。分配 tensor memory 需要执行:
# allocate all 2^18 bytes of smem for tcgen05, all 512 cols allocated if elect_one_warp: tcgen05_alloc(ptr_tmem_addr, max_tmem_cols) # Ensure all threads see initialized mbarrier and # tensor memory allocation barrier() tmem_addr = ptr_tmem_addr[0]这个分配相当不平凡:首先,分配必须由单个 warp(而非单线程)发起;其次,必须绕道共享内存才能拿到分配好的tmem地址。
tcgen05.mma的输入与配置被编码进描述符(descriptor):
- 指令描述符(
idesc):编码指令形状、数据类型、矩阵布局等。由于这些属性在计算过程中保持不变,该描述符在迭代中不变。 - 共享内存描述符(
adesc、bdesc):编码矩阵 A、B 的共享内存布局与访问模式。由于遍历矩阵不同 K 切片时共享内存地址会变化,这些描述符会在num_k_mma次迭代中被递增。
更深入的解释见本文附录。
MMA 完成后会到达mma barrier。其工作原理与tma barrier基本一致:阻塞所有线程直到 MMA 完成。这样,在当前 tile 上的所有 MMA 完成之前,不会有线程进入下一轮迭代去发起 TMA 操作。
TMEM → 寄存器
至此我们已覆盖两个主要函数:
for i in range(K // BK): load_tiles_ab() # leader thread loads A and B tiles issue_mma_axb() # leader thread issues MMA(A x B)结果已累加并存储在 tensor memory 中。下一个问题是:如何把它从 tensor memory 搬进全局内存?
唯一能把数据搬出 tensor memory 的方式是先把数据搬进寄存器。这通过tcgen05_ld操作完成:
c_frag = tcgen05_ld datapaths=16, bits=256, repeat = BN // 8, dtype=accum_type, pack=False, width=c_frag_size, tcgen05_load_wait() # wait for the load to finish这条指令相当复杂,逐步拆解:查看 tensor memory 中数据的存储方式,可以发现 tensor memory 存放一个64x64的C_tile。其布局组织与访问模式(依据 NVIDIA Parallel Thread Execution ISA 9.0 中 tcgen05 数据路径布局)如下:
因此,要访问这块内存,块内每个 warp 需要读出 16 个 lane,整个 warp-group(4 个 warp)读出 64 个 lane。参数(datapaths和bits)正是用来指定这个加载模式,tcgen05_ld内部派发tcgen05.ld.16x256b指令来加载每组 lane。
这意味着每次迭代,线程沿 tensor memory 的列方向加载 256 bits,即 8 个元素(不是 16 个——记得第一篇文章中我们以 FP32 累加结果以保精度,每个元素占 4 字节),共BN/8次迭代。于是 warp 内 32 个线程中的每一个必须持有 4 个元素。
重复BN//8 = 8次后,每个线程在一个寄存器数组中持有 tile 的 32 个元素。确认所有数据都成功转移到寄存器后,就可以释放之前分配的 tensor memory:
if elect_one_warp: tcgen05_release_allocation_lock[1]() tcgen05_dealloc1寄存器 → 全局内存
目前的代码进展:
setup_kernel() for i in range(K // BK): load_tiles_ab() # leader thread loads A and B tiles issue_mma_axb() # leader thread issues MMA(A x B) transfer_c_tile_to_registers() # move final C tile from tmem to registers还缺关键一步:write_c_tile_to_global_memory——把数据从寄存器搬进全局内存。先确定要写到哪里:矩阵是 4096x4096,每个 block 负责输出其中的一块64x64tile。以block_idx.y = 2, block_idx.x = 2为例,它负责输出第 3 行的第 3 块 tile:
用LayoutTensor.tile()方法提取输出矩阵的一块 tile:
ctile = c.tileBM, BN再为每个 warp 进一步分块:
c_gmem_warp_tile = ctile.tileBM // num_warps, BN聚焦 warp 0 的 tile:c_gmem_warp_tile的 tile 0 表示前 16 行 x 64 列(16xBN),需要把这个 16x64 的 tile 映射到 warp 0,因为累加值就在那里。下图展示了tcgen05.ld.16x256PTX 指令中元素到 lane(线程)的映射:
这里涉及相当多的索引计算。有没有办法在 warp 的 tile 上创建视图——"小口袋"——让每个线程精确拿到自己需要写入数据的布局?Mojo 正好提供了这样的库函数,可以简洁地完成:
c_gmem_frag = c_gmem_warp_tile.vectorize[1, 2]().distribute Layout.row_major(8, 4) )对刚接触LayoutTensor的读者可能有点复杂,可视化 thread 0 的视图:代码第一部分意识到,由于每个线程存储 2 个连续元素,16x64 的 tile 可以看作 16x32 的"2 值向量" tile:
随后.distribute[Layout.row_major(8, 4)]把这个 16x32 的向量分布到 8x4 个线程上,循环往复:
偏移按row_major(8, 4)(lane_id())计算。例如 thread 0 在所有子矩阵中取(0, 0)处的向量(图中绿色格子),thread 6 取(1, 3)(图中蓝色格子)。事实上,每个子矩阵与 NVIDIA 的 Figure 185 布局完全一致。最终得到2x8个子矩阵,每个子矩阵存放8x4个 2 值向量。distribute正如我们所承诺的,给了每个线程它所需"口袋"的视图:
有了这个映射,向全局内存输出就只是一个平凡的循环:
alias num_vecs_m = c_gmem_frag.shape[0]() alias num_vecs_n = c_gmem_frag.shape[1]() comptime for n_vec in range(num_vecs_n): comptime for m_vec in range(num_vecs_m): alias i_vec = n_vec * num_vecs_m + m_vec c_gmem_frag[m_vec, n_vec] = [c_frag[2 * i_vec], c_frag[2 * i_vec + 1]]以num_vecs_n, num_vecs_m = (8, 2)为例:跨每个 warp,一次写出一个子矩阵——先沿 M 维度写 2 次,再沿 N 维度写 8 次。循环执行过程如下:
以上是针对单个 warp 的。放大到 CTA 级别,可以把 CTA 的 tile 映射到全局内存中的C矩阵:
为上述一切配置共享内存
先看看 SM 上共享内存栈长什么样,这正是此前跳过的setup()阶段。共享内存主要用于:输入 tile、内存屏障和 TMEM 分配。
var a_smem = external_memory[Scalar[a_type], address_space = AddressSpace.SHARED]()) # Offset BMxBK for A tile var b_smem = (a_smem + a_size).bitcast[Scalar[b_type]]() # Offset BNxBK for B tile var tma_mbar = (b_smem + b_size).bitcast[Int64]() # Offset 8B for tma memory barrier mma_mbar = tma_mbar + 1 # Offset 8B for mma memory barrier ptr_tmem_addr = mma_mbar + 1上面的设置代码从动态共享内存分配(external_memory)拿到基地址,然后按下图方式逐步增加偏移:
把各部分拼起来并基准测试,这个 kernel 达到155.0 TFLOPS——比朴素 kernel 提升了28 倍。但换个角度看,它仍然只有 cuBLAS 性能的8.7%:
Kernel 3:Swizzling
Kernel 2 的一个开销是加载输入 tile 时需要发起多次 TMA 调用。原因在于BK=64,而 Tensor Core 需要的规范化布局只允许每次按 K 复制 16B。还有其它支持更大 K 维的布局——例如最宽的128B 布局。数学计算表明,只要配合Swizzle<3, 4, 3>,我们确实可以用单个行主序BM x BK(BK=64)tile。什么是 swizzle?为什么是<3, 4, 3>这个神奇组合?要理解它,先温习一下共享内存。
共享内存的银行(banks)
共享内存由 32 个连续的、4B 宽的 bank 组成:
共享内存中每个 bank 每周期只能服务一次请求,而访问不同 bank 的多个线程可以在同一周期内被服务。也就是说,bank 0服务thread 0、bank 16服务thread 1,可以同时进行:
银行冲突(Bank Conflicts)
但如果两个请求访问同一个 bank 呢?例如两个线程访问bank 0的不同地址——比如 thread 1 现在想访问row 3 column 0的元素:
这需要 2 个周期:bank 0先服务thread 0,一个周期后,bank 0(图中表述为 bank 2 服务 thread 1)再服务thread 1。直观上也讲得通:为了最大化吞吐,GPU 被设计为每个周期扫过所有 bank(32 个 bank x 每个 bank 4B)最多加载 128B;同一 bank 的第二次加载只能排到后面的周期。
注意指令是由 warp 发起的。当 warp 内线程访问映射到同一 bank 的不同地址时,硬件不得不把执行拆成多个周期。这种执行停顿就叫bank conflict,显然对性能有害。
把这个规律套到 128B 规范化布局(tile 为BM x BK且BK=64):第一个 core matrix 的 8 行全部映射到相同的 bank0-3:
这会给每个 core matrix 制造 8 路 bank conflict,导致每行的写入串行执行。显然需要一种技术,在读取所需数据时不产生这些停顿。
Swizzling 原理
Swizzling 就是解决 bank conflict 的技术:用按位异或(^)交换索引,让数据不再落在同一个 bank。用一个例子演示——为简单起见假设有 16 个 bank:
注意:不同行上的相同索引(1-16)已被交换到不同的 bank——也就是说,当线程按相同索引访问不同行的元素时,不再发生 bank conflict。
128 字节 Swizzling
解读 128B swizzle 模式<3, 4, 3>:第一个3对应2^3 = 8——core matrix 的行数;4对应2^4 = 16B——core matrix 的宽度(8 个元素 x 2B);最后一个3是2^3 = 8,意味着 8 个 16B 的块横跨全部 32 个 bank(128B)。有了这些值,swizzle 函数就提供了正确的 XOR 模式来解决 core matrix 的 bank conflict。这个模式可以在 Mojo 中直接写出来(参见仓库中的 swizzle.mojo,Swizzlefunctor 的实现),并针对常见模式做了泛化。可视化如下:
每 8 个元素(16B = 8*2B)通过清零xor操作数中的 3 个最低有效位来分组。这个xor计算就像前面演示的那样,在每一行内部交换分组。结果就是每 8 个元素分布在不同的 bank 上,并像下图那样延续:
完整的数学细节见附录。看看两个相邻 core matrix 是如何被 swizzle 到 32 个 bank 上的:
加入 swizzle 后变为:
同一个 core matrix 中任意两个元素永远不会落在同一个 bank——因为 core matrix 宽度为 16 字节,因此配合 128 字节 swizzle 不存在 bank conflict。这正是 swizzling 极其有用的原因,每个高性能 GPU kernel 都会使用它。
更新后的内核
代码改动极小,因为对 swizzling 的支持来自库的 layout tensor 与指令本身。唯一要改的是告诉 TMA 和tcgen05.mma采用哪种 swizzle 模式:
alias a_swizzle = TensorMapSwizzle.SWIZZLE_128B alias b_swizzle = TensorMapSwizzle.SWIZZLE_128B #for the tma, used on writing in data from global memory alias a_smem_layout = tile_layout_k_major[ a_type, BM, BK, swizzle_mode=a_swizzle ]() alias b_smem_layout = tile_layout_k_major[ b_type, BN, BK, swizzle_mode=b_swizzle ]() #for the mma adesc = MMASmemDescriptor.createaSBO, aLBO, a_swizzle bdesc = MMASmemDescriptor.createbSBO, bLBO, b_swizzle由于LayoutTensor理解 swizzling,我们可以把 swizzle 操作的细节隐藏在 layout tensor API 背后,其余代码保持不变。仓库中 swizzle.mojo 的Swizzlefunctor(位于max/kernels/src/layout/目录)即是对该模式的底层实现:构造时根据bits、base、shift计算yyy_mask与zzz_mask,调用时执行offset ^ shiftr(offset & self.yyy_mask, self.shift)。
性能
经过上述优化,我们在 B200 上达到288.3 TFLOPS(87% 的提升)。换句话说,共享内存 bank conflict 的影响几乎把性能砍掉了一半;解决 bank conflict 后,我们达到了 cuBLAS 的16.4%,正在快速缩小差距:
Kernel 4:在共享内存中打包输出并利用 TMA Store
前一个 kernel 的输出每次向全局内存写两个连续的 BF16 值——每次 store 只有 4B,而 Blackwell 单条 store 指令(st.global.v8.b32)最多支持 32B。此外,我们还可以用 TMA 每条指令 store 整个输出 tile,减少发出的指令数。
在共享内存中打包输出
要利用 TMA store,需要先把输出数据打包进共享内存:在把输出从 tensor memory 加载到寄存器之前,先把寄存器复制到共享内存。由于全局内存中的输出是 BF16,必须在复制到共享内存前把寄存器从 FP32 转成 BF16。
但输出结果在寄存器中是按特定布局(16x256bits 加载,见上文 TMEM→寄存器一节)分片的,因此把寄存器复制到共享内存时需要处理好这一点。幸运的是,NVIDIA 提供了stmatrix指令:它以精确的 16x256 bits 布局,把8x16B的 core matrix 分布存储到共享内存,并且允许用户为每一行指定共享内存中的地址。256 bits(32B)与每行 16B 之间存在明显的不匹配——因为从 TMEM 加载的数据是 FP32,存入共享内存时转为 BF16。stmatrix每条指令最多存储 4 个 core matrix(2x2)。因此,打包16x64(BN=64)的 warp tile 需要 4 次stmatrix迭代:
注意,这个操作在写入共享内存时同样会遇到 bank conflict 问题,因此我们使用 128B swizzling(BN * 2B = 128B)来避免冲突。
TMA Store
数据在共享内存中完成 swizzle 和打包后,就可以发起 TMA store 操作,把数据异步复制回全局内存。下面的代码展示了 TMA store 及其同步方式。在 TMA 发起异步 store 之前,需要通过fence_async_view_proxy做内存 fence,确保之前的共享内存打包结果对 TMA store 可见:
# Launch one TMA store per thread if elect_one_warp and thread_idx.x < BN // TMA_BN: # memory fence to ensure previous shared memory access # is seen by TMA instruction fence_async_view_proxy() c_tma_tile = ... # setup the tile for tma # c_tma_op is created similarly like a_tma_op for loading data c_tma_op.async_store( c_tma_tile, (block_idx.x * BN + thread_idx.x * TMA_BN, block_idx.y * BM), ) # Commit TMA store c_tma_op.commit_group() # wait for the store to complete c_tma_op.wait_group[0]()发出 TMA store 后,先用commit_group()提交这些 store——它把从上一次 commit 到当前程序计数器之间发出的 store 归为一组。随后的wait_group[N]()会等待直到只剩N组 store 还在传输中。例如,若有 3 个已提交的组,wait_group[2]()确保第一组完成、后两组仍在传输。上述代码中的wait_group[0]()守护所有 TMA store 完成。按 commit group 等待的能力允许你构建流水线,并在后续优化中高效地重叠其它任务。
TMA store 与 TMA load 还有一个区别:多个线程可以并行发起 TMA store:
# Launch one TMA store per thread if elect_one_warp and thread_idx.x < BN // TMA_BN:这里TMA_BN取决于 swizzle 模式(例如 128B swizzle + BF16 时为TMA_BN=64)。如果 tile 维度BN更大,就需要把维度除以TMA_BN并发起多个 TMA store。例如BN=128对应两次 store,由两个线程发起以最大化并行度:
性能与剖析
这个 kernel 的性能基本持平,为293.6 TFLOPS(准确说是慢了 0.7%)。为什么?因为性能从根本上仍受限于全局内存访问:
下图是来自 NCU 的计算与内存吞吐剖析:绿色柱是 kernel 3,蓝色柱是 kernel 4。可以看到两个 kernel 的计算与内存吞吐都很低:
此外,TMA store 的真正威力在于其异步性——它开启了流水线与操作重叠的可能性。当前 kernel 为我们在后续文章中利用这些特性打好了基础。
总结:本文演示了如何对 matmul 做分块,以及如何用 TMA load/store、tcgen05.mma、stmatrix等特性,以最优指令集编程 Blackwell GPU。这一系列努力带来了相对朴素 kernel58 倍的提升,但仍落后于 cuBLAS 的性能。
后续文章将在本 kernel 基础上进一步优化底层的调度与执行算法:下一篇将展示如何构建 warp 专用流水线,重叠数据传输与计算,以获得更接近业界最先进的性能。
附录
描述符:LBO 与 SBO
tcgen05.mma用描述符指定输入数据在共享内存中的布局以及指令形状、数据类型等。在 Mojo 中创建 smem 描述符:
adesc = MMASmemDescriptor.createaSBO, aLBO, a_swizzleMMASmemDescriptor负责以tcgen05.mma要求的格式编码所有这些信息。其中最重要的细节是LBO和SBO:
- LBO(leading dimension byte offset,前导维字节偏移):K 维度上两个相邻 core matrix 之间的字节数。
- SBO(stride dimension byte offset,步进维字节偏移):
M/N维度上两个相邻 core matrix 之间的字节数。
在 kernel 2(无 swizzling)中,对 A 打印结果是:
aSBO=128 aLBO=1024如下图所示,LBO为 1024B,因为两列 core matrix 之间的距离是BM*16B = 1024B;SBO为 128B,因为每个 core matrix 的大小是8x16B = 128B:
UMMA 描述符idesc的模式类似,只是它是 32 位的,并额外编码了稀疏性、数据类型、矩阵是否转置等信息。其详细编码可参考仓库中的实现(max/gpu/compute/arch/mma_nvidia_sm100模块及 mma.mojo)。
Swizzling 数学
swizzling 的数学定义如下。给定定义为Swizzle(bits, base, shift)的 swizzle:
# 0bxxxYYYxxxxZZZxxxx # ^--^ Base is the number of least-sig bits to keep constant # ^-^ ^-^ Bits is the number of bits in the mask # ^------^ Shift is the distance to shift the YYY mask 1) ZZZ is the first mask, extracted right after the base 2) YYY is the second mask, extracted shift after the base 3) We XOR these two, to get AAA=YYY XOR ZZZ 4) We place this new substring in place of the first mask, ZZZ 5) Final answer becomes: # 0bxxxYYYxxxxAAAxxxx在仓库的 Mojo 底层代码 swizzle.mojo 中,swizzle 实现为:
bit_msk = (1 << bits) - 1 self.yyy_mask = bit_msk << (base + max(0, shift)) self.zzz_mask = bit_msk << (base - min(0, shift)) swizzled = offset ^ (offset & self.yyy_mask) >> shift考虑 128B swizzle:bits=3、base=4、shift=3。数学上意味着取输入地址的 7-9 位作为掩码,与 4-6 位做xor,从而生成 kernel 3 中展示的模式:
仓库中的 max/kernels/src/layout/swizzle.mojo 提供了完整的Swizzlefunctor 实现,并在max/kernels/src/linalg/matmul/gpu/sm100_structured/default/matmul_kernels.mojo等生产级 SM100 matmul kernel 中得到实际应用——这正是本文所述优化技术落地为真实推理/训练内核的例证。
延伸阅读:本系列其余文章见 matmul-on-blackwell-part-1.md(Blackwell 架构与朴素 kernel)与 matmul-on-blackwell-part-3.md(warp 专用流水线优化)。
【免费下载链接】mojoThe Modular Platform (includes MAX & Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考