导读:从朴素内核到 cuBLAS 的 94%——CUDA 矩阵乘法优化工作记录

4 minute read

Published:

这是一篇导读,不是转载。 原文是 Simon Boehm 于 2022 年 12 月发表在个人博客上的工作记录 How to Optimize a CUDA Matmul Kernel for cuBLAS-like Performance: a Worklog, 约 9300 词、18 个代码块。作者自述这篇写得比平时”粗糙”——它更像一本记满想法和草稿的笔记本, 记录了他如何一版一版地把矩阵乘法内核推到接近 cuBLAS 的水平。

本文用自己的话梳理它的技术脉络,补充我对其中方法的理解与延伸。正文不含原文段落,也未搬运原文插图; 所有代码、数据与完整推导请以原文为准(原文提供全部内核源码)。

为什么这篇值得花时间

矩阵乘法几乎构成了大模型训练与推理中的全部浮点运算量。这意味着这个内核的性能上限,在很大程度上就是一个硬件平台的有效算力上限——你写的注意力、MoE 路由、卷积,最终都要落到 GEMM 上。

但绝大多数人使用的 GEMM 是别人写好的(cuBLAS、CUTLASS、Triton 生成的内核)。我们知道自己”在调库”,却说不清这 100% 的性能是怎么来的、还差多少、瓶颈在哪一层。

这篇工作记录的价值在于它把这个黑箱逐步拆开:作者从最朴素的实现出发,每加一个机制就测一次,把每一版的收益、代价和瓶颈都记录下来。九个版本走完,性能从 cuBLAS 的 1.3% 提到 93.7%。这个过程本身就是一份”GPU 性能模型”的实操教材。

实验设定与两个上界

作者的目标不是做一个 cuBLAS 替代品,而是理解硬件的性能特征。所以他的做法很有代表性:先算上界,再动手写代码。

硬件是一块 NVIDIA RTX A6000(Ampere,计算能力 8.6),关键参数我列在这里,因为后面每一步优化都在和这些数字打交道:

项目数值
SM 数量84
每 SM 最大常驻 warp 数48
每 SM 最大线程数1536
每 block 共享内存上限48 KB(每 SM 共 100 KB,可与 L1 互相划分)
fp32 峰值算力~30 TFLOPs/s
全局内存带宽~768 GB/s

问题规模是两个 4092×4092 的 fp32 矩阵相乘再加一个同尺寸矩阵(即 SGEMM 的 C = αAB + βC)。

然后是那段我认为全文最该学的”餐巾纸演算”:

  • 总计算量:4092² 个输出元素,每个都要做长度为 4092 的点积。乘加被硬件映射成一条 FMA 指令,但仍计两次浮点运算,所以总量是 2 × 4092³ + 4092² ≈ 137 GFLOP。
  • 最小访存量:读入三个矩阵 3 × 4092² × 4B ≈ 201 MB,写出一个 4092² × 4B ≈ 67 MB,合计约 268 MB(作者在同一段里也写作 278 MB,是取了整)。
  • 计算耗时下界:137 ÷ 30000 ≈ 4.5 ms。
  • 访存耗时下界:268MB ÷ 768GB/s ≈ 0.34 ms。

两个下界差了一个数量级(约 13 倍)。这条比值就是设计的靶心:只要你的内核访存量控制在这个最小值的十倍以内,计算就是瓶颈,访存就不是——你应该把所有注意力放在怎么把 FMA 流水线喂满。

作者顺带给了两个很有信息量的参照:cuBLAS 实际搬了约 500 MB(约为理论最小值的 1.9 倍,说明它做了相当好的分块复用);而如果换成 TF32/BF16 让 cuBLAS 用上张量核心,同一次计算只要 0.44 ms——这会儿访存反而成了瓶颈。这个对比很直观地解释了为什么低精度推理的优化重点会整个换一个方向。

九个版本:每一版解决了什么

作者最终留下的内核有九个(第 7、8 版是解决共享内存 bank conflict 的中间产物,虽然消除了冲突但整体更慢,被略过)。下表是原文的性能汇总,我把每一版的机制也一并标注出来:

版本关键机制GFLOPs/s相对 cuBLAS
1朴素实现:一线程一输出,直接点积309.01.3%
2全局内存合并访问(coalescing)1986.58.5%
3共享内存分块缓存(cache blocking)2980.312.8%
4一维 blocktiling:每线程算多个结果8474.736.5%
5二维 blocktiling:提升算术强度15971.768.7%
6向量化访存(float4 / LDS.128)18237.378.4%
9自动调参(5 个模板参数)19721.084.8%
10Warptiling21779.393.7%
cuBLAS(对照组)23249.6100%

下面是我对这九步的理解。

第 1 版:朴素实现——先建立一个”错的基准”

思路最直接:grid/block/thread 三层结构铺满输出矩阵 C,每个线程负责一个元素,用点积算出结果。因为 C 的每个位置只被一个线程写,不需要任何同步。

它跑完一次要 0.5 秒左右,约 300 GFLOPs——和 2015 年那块 Haswell CPU 上优化过的 BLAS 库差不多。这个数字很有冲击力:一块 2020 年的旗舰级 GPU,用最直接的方式写 GEMM,性能等于五年前的 CPU 库。

失败的原因可以算出来。同一 block 内相邻线程取 B 的同一列、A 的不同行,在最坏情况下(假设零缓存命中),每个线程要从全局内存读约 2 × 4092 个 float。乘上线程总数就是 548 GB 的访存——相当于理论最小值 268 MB 的两千倍出头(548GB ÷ 268MB ≈ 2045)。带宽再高也扛不住这种放大。

这里还顺带出现了一个概念,作者叫它 tile quantization:输出尺寸不能被 block 尺寸整除时,最后必须多启动一批 block,而其中大部分线程是闲置的。这是把”固定大小的分块”映射到”可变大小的输入”时必然出现的浪费。

第 2 版:合并访存——在正确的粒度上思考

要理解这一步,绕不开 warp 这个概念:block 内的线程在执行时按 32 个一组被划分成 warp,warp 是分配给 warp scheduler 的调度单位。

关键细节是 warp 的划分依据——它按 threadId 连续分组,而多维 block 的 threadId 是这样算的:

threadId = threadIdx.x + blockDim.x * (threadIdx.y + blockDim.y * threadIdx.z)

也就是说 x 维是在 warp 内连续的那一维。这解释了一个初学者很容易踩的坑:直觉上(行主序的世界里)我们觉得 x 是”行”、y 是”列”,但在 warp 的视角下,x 才是”快变”的那一维。

明白了这一点,合并访存就顺理成章了:同一 warp 内的线程如果访问连续地址,这些请求可以被合并成更少的访存事务。第 2 版正是让相邻线程去写 C 中相邻的地址(也就是让线程沿着列方向铺开),访存效率立刻上了一个台阶——从 309 提到 1986 GFLOPs,6.4 倍。

第 3 版:共享内存缓存——把复用的数据搬到片上

全局内存很大但在片外,共享内存(SMEM)很小但在片上,延迟和带宽差着数量级。作者引用的 Volta 实测数据是全局内存约 750 GiB/s、共享内存约 12000 GiB/s——十六倍左右的差距

于是策略很清楚:让一个 block 先把 A、B 的对应分块协作搬进 SMEM,然后 block 内所有线程从 SMEM 反复取用。这在概念上就是把”每个元素都去片外拿一次”变成”每个分块只去片外拿一次,之后在片上复用”。

代价随即出现,作者专门算了一遍 occupancy(每 SM 实际常驻 warp 数 / 最大可常驻 warp 数):这个版本的 SMEM 用量是每 block 8 KB、每线程 37 个寄存器,算下来大约 66% 的占用率。

这里有个反直觉的结论值得记住:66% 的占用率并不算问题。作者指出,无论算术强度极高还是极低,达到峰值吞吐其实都不需要高占用率,中间那段”不上不下”的区域才需要(这正是 Volkov 讨论的 cusp behavior)。所以在优化时看到占用率不高,不要条件反射地去调它。

第 4 版:一维 blocktiling——让每个线程多算几个结果

到这里瓶颈已经不是”访问模式差”,而是”每次访存的复用次数不够”。解法是让每个线程负责多个输出(沿一维展开),把结果累加在寄存器里。

这一步有一个很值得看的编译器细节(原文的 Sidenote)。作者最初的写法手动把 B 的某一项缓存进局部变量、并调整了内层两个循环的顺序,看起来省下了不少 SMEM 访问。但实测发现,如果不做这个改写,性能也没有变差。原因要去看 SASS:编译器其实已经根据循环结构自己完成了等价的优化,生成的指令流差异很小。

这件事的方法论意义大于技术意义:在改代码之前先看汇编。源码层面的”明显优化”经常是无效的,真正的依据应该是 SASS/PTX。

效果很显著——8474.7 GFLOPs,相对上一版的 2980.3 提升了约 2.8 倍。

第 5 版:二维 blocktiling——算术强度的正面进攻

作者在这里给出了全文我最认同的一句判断:算术强度(arithmetic intensity,即每次 GMEM↔SMEM 传输所支撑的 FLOP 数)是唯一值得优化的目标函数。

为什么”算更多的结果”能提高算术强度?因为输入可以被更多输出共享。算一个 8×8 的结果方阵,比算一个 8×1 的列,能摊薄更多输入加载。

这一版把每线程的输出扩成 8×8,性能直接到 15971.7 GFLOPs(约 16 TFLOPs)。作者同时指出,到这一步 SMEM 访问次数已经从”每个结果 K 次”降到”K/4 次”,GMEM 访问降到”K/64 次”,但因为内存流水线拥塞导致的 warp stall 仍然偏多——说明还有空间。

顺带说,作者整篇反复强调”继续优化算术强度,直到不再是内存受限为止”,这条判据比”占用率”、”指令数”之类的指标都更可靠。

第 6 版:向量化访存——用宽度换效率

两个动作:

一是把 SMEM 里的 A 转置存放,这样读 A 时也能用上 128 位的向量加载(SASS 里的 LDS.128),而不是 32 位逐元素加载。单这一项带来约 500 GFLOPs、3% 的提升。

二是把 GMEM 的读写全部换成 float4 向量类型。这里可以看到加载 A 的同时顺手完成转置的写法——一次 128 位读取,拆成四个分量写进转置后的位置。

这一步到 18237.3 GFLOPs(78.4%)。到这里,一个内核该有的基本机制就齐了:分块、复用、合并、向量化。

第 9 版:自动调参——承认手调不动了

到这一版,内核暴露出了五个模板参数:BMBNBK(控制 GMEM→SMEM 的分块大小)和 TMTN(控制 SMEM→寄存器的分块大小)。第 6 版用的是 BM=BN=128BK=TM=TN=8

作者的选择很务实:写脚本暴力搜。他从所有组合里筛掉不合理的配置(比如为了用上向量化 SMEM 加载,BM*BK 必须能被 4*NUM_THREADS 整除),剩下约 400 种,逐一基准测试。

他诚实地说了一句很关键的话:最优参数在不同 GPU 型号上有明显差异。这也解释了为什么 Triton 这类编译器要内置自动调参流程——而 cuBLAS 大概率是维护了一张”GPU 型号 → 预计算配置”的映射表。

收益:19721.0 GFLOPs,84.8%。

第 10 版:Warptiling——把最后一级并行显式化

最后一步是在 block 分块和线程分块之间,再插入一级 warp 分块

Warp 这个层级有个认知门槛:它在 CUDA 代码里根本不出现,是纯硬件概念,没有对应的软件抽象。但它在性能上极其关键,原因至少有三个——warp 是映射到 warp scheduler 的调度单位;SMEM 的 bank conflict 只发生在同一 warp 的线程之间;较新的 GPU 上有寄存器缓存,更紧的线程分块能提高寄存器局部性。

完成这一版之后,三级并行结构就完全显式了:block 分块对应跨 SM 并行、warp 分块对应跨 warp scheduler 并行、线程分块对应指令级并行。

最终成绩 21779.3 GFLOPs,93.7%

作者的时间账单

我觉得这个细节值得单独拎出来,因为它比性能数字更接近真实的工程感受:

前六个内核(达到峰值 FLOPs 的 80%)花了他两个周末;后面的自动调参和 warptiling(从 80% 推到 94%)又花了四个周末

而且他坦然承认,剩下那 6% 他暂时不追了——因为”写这段代码带给我的学习收益也在递减”。这是个很健康的判断。在真实工程里,知道该在哪里停下,和知道该优化什么同样重要

作者还给了一个未完待续的清单(第 11 版的方向):双重缓冲(在 GMEM→SMEM 和 SMEM→寄存器两级做流水重叠)、Hopper 上引入的 warp specialization 与直接 GMEM→SMEM 加载指令(用来降低寄存器压力)、彻底消除 SMEM bank conflict、以及通过阅读 Triton 生成的 PTX 来理解现代自动生成内核的做法。

我从中带走的方法论

抛开具体的内核技巧,我认为这篇记录里可迁移的东西有这么几条:

先算上界,再动手。 四行餐巾纸演算就确定了两件事:这个问题的本质是计算受限;以及访存量只要控制在最小值的十倍(约 2.7 GB)以内就不会成为瓶颈。没有这个上界,你无法判断一次优化是”接近上限了”还是”还有十倍空间”。

每次只改一个变量,并且用 profiler 验证。 作者全程在看 warp stall、occupancy、指令混合比,而不是靠直觉。整个过程的叙事也是这样推进的:测量 → 发现瓶颈 → 引入一个机制 → 再测量。

算术强度是最可靠的目标函数。 相比占用率、指令数这些容易被误读的指标,”每次数据传输支撑多少计算”直接对应了 roofline 模型里工作点向右移动的方向。

改代码前先看汇编。 第 4 版那个 sidenote 是个很好的反面教材——源码层面看起来明显的优化,编译器早就做了。

把可视化当成写作的一部分。 作者说了一句让我印象深刻的话:一旦把内核该长什么样的图画清楚,代码写起来出奇地容易。这类优化工作里,真正的难点通常不在编码,而在把数据流动想明白。

与今天的工作流的对照

这篇是 2022 年的记录,而今天大多数人写 GEMM 用的是 Triton、CUTLASS 或者直接 cuBLAS。但这篇的价值并没有因此下降,反而更清晰了:

上层框架替你做的事,正是作者手工做的这些。 Triton 的 BLOCK_M/BLOCK_N/BLOCK_K 就是 BM/BN/BK;它的自动调参就是在做第 9 版那件苦力活;tl.load 的向量化由编译器决定;bank conflict 和缓存策略由后端 pass 处理。而 CUTLASS 更直接——它的分层抽象(threadblock / warp / thread)就是第 10 版的显式化,pipeline 抽象就是第 11 版想做的双重缓冲。

所以当你需要诊断一个 GEMM 为什么慢的时候,能定位到”是没有合并访存、还是算术强度不够、还是 SMEM bank conflict、还是占用率撞上了 cusp 区”,判断依据恰恰来自这篇里手工做过一遍的经验。库把优化自动化了,但没有把性能模型自动化。

如果你在写 CUDA/Triton 内核、或者在做推理引擎的性能调优,我建议按这个顺序读原文:先看开头的餐巾纸上界和性能汇总表,再挑其中两三个版本细读(第 2 版的 warp/合并访存、第 5 版的算术强度、第 10 版的 warptiling 是我认为信息密度最高的三段),其余当作查阅。

原文与相关材料

如果你想先看这篇里几个论点的具体实现,我建议直接从上面那个源码仓库入手,对照第 2、5、10 版的内核读——那三版分别对应”访问模式”、”数据复用”和”并行层级显式化”三个不同层面的问题。