我在做推理优化的时候,用性能分析工具看了下整个计算图,发现一个很扎心的事实:一个普通的矩阵乘法算子,就能吃掉单次迭代接近四成的时间。当时第一反应是换参数、调库、换格式,折腾一圈之后发现,通用数学库在那个特定形状下就是上不去,不管怎么切都是中等偏上的表现。那段时间正好在调研算子层自研的可行性,于是动手写了一个自己的矩阵乘内核项目,代号 DeepGEMM。这篇文章不打算讲太多高深理论,而是想把这个项目从“为什么做”到“怎么写”,再到“实测怎么调优”的过程完整摊开,给同样想在算子层动手的读者一条能直接参考的路。
DeepGEMM 的目标其实非常朴素:在固定形状集合内,做到比厂商通用库快 10% 以上;支持 FP16、BF16、TF32 混合精度;接口简单,能方便地嵌入到前向推理或训练脚本里。后面我会从需求分析、计算核心拆解、流水线调度、实测调优和后续规划几个部分逐一展开,整个过程基本就是我实际做这个项目时的思考顺序。
1. 为什么动手写一个 GEMM 内核:一个不算“划算”的决定
1.1 一次性能剖析引发的思考:通用库并不是万能的
先说那次 profiling。当时跑的是一个 batch 不大、序列长度偏长的模型,热点集中在 QK^T 和注意力输出的投影矩阵乘上。矩阵形状大概是 M = 2048、N = 512、K = 1024,这个形状不算刁钻,但通用数学库给出的性能只有该硬件理论峰值的五成左右。问题在于它是一个整体算子,库内部有很复杂的启发式分派逻辑,会选择它认为最优的算法,但这些启发式规则更多是面向“大多数形状”调优的,遇到某些偏窄或偏长的矩阵组合,未必能挑到真正适合的 kernel。
更让人难受的是黑盒。想给这个热点算子加一个偏置、加一个缩放,就得把矩阵乘的结果先从显存过一遍再做下一次操作,白白增加读写开销。如果内核是自己写的,可以在矩阵乘结束之后顺手把 bias、scale、激活函数全做掉,省一次甚至两次全局内存往返。
“别重复造轮子”这句话我当然认同,但它有个前提:通用轮子能在你的场景里跑得足够好。当通用库在关键形状上不给力,而你又对延迟和吞吐有持续优化需求时,自研就不是为了炫技,而是为了可控。所以这里也想给读者一个判断标准,什么情况下值得自己动手:
- 你的工作负载里,热点矩阵乘的形状足够固定,不会三天两头变化;
- 有明确的性能瓶颈,且厂商库无法满足;
- 你对目标硬件架构有基本的了解,至少能读懂性能分析工具给出的指标;
- 愿意投入一个较长的迭代周期去调优,而不是指望一次写出完美内核。
1.2 DeepGEMM 的项目定位:只在选定的战场上打赢
我不打算做一个什么都支持、什么形状都能跑的“万能矩阵库”。那工程量太大,也不是一个人短期能做出来的事。DeepGEMM 的定位是:只针对一组固定形状集合,做深做透。比如先锁定 M 在 1024 到 8192 之间、N 在 512 到 4096 之间、K 在 512 到 4096 之间的组合,且以混合精度为主。在这个范围里,我可以放心大胆地手工选择 tile 大小、寄存器分配和流水线级数,而不是像通用库那样做一个“八面玲珑”的调度器。
这也直接影响后续的设计取舍。因为形状固定,我可以按 M、N、K 预先切分好 block,不需要在运行时去做复杂的启发式搜索。因为精度固定,我可以针对 Tensor Core 指令专门设计共享内存布局,不需要同时兼顾各种不支持的路径。这种“自我限制”反而让内核可以做得更激进,也更容易达到目标。后面的实测数据也证明了这一点:当厂商库在某形状上只能跑 620 TFLOPS 左右时,DeepGEMM 可以在相同硬件上做到接近 720 TFLOPS,领先约 16%。
这里想多说一句:自研算子不是“实验室玩具”,在生成式模型、大规模推荐系统这类追求极致成本效率的场景里,算子是真正值得一帧一帧去抠的地方。
2. DeepGEMM 的计算核心拆解:从数学公式到 GPU 内核
2.1 从算式到分块:为什么不能把整个矩阵一把梭
矩阵乘法本身很简单:C[M, N] = A[M, K] * B[K, N]。核心计算量是 2 * M * N * K 次浮点操作。比如 M=N=K=4096 时,单次矩阵乘就是大约 137 GFLOP 的计算量。这个数字听起来很大,但现代加速卡的算力动辄每秒上百万亿甚至更多,所以瓶颈往往不在“能不能算”,而在“能不能把数据喂到计算单元里”。
一个关键约束是:片上存储非常有限。一个计算单元里的共享内存通常只有几十到两百多 KB,寄存器总量也有限,而 A 和 B 矩阵往往有几十 MB 甚至更大,根本不可能一次性放入片上。于是就有了分块矩阵乘的思路:把输出矩阵 C 切成若干个 BM x BN 的小块,每个线程块负责计算一个小块。在 K 维度上,每次只加载一段长度为 KK 的数据到共享内存,计算完再加载下一段。
我举个例子。如果选定 BM = 128、BN = 128,那么每个线程块负责输出一个 128x128 的矩阵块。A 矩阵需要取 128 行,B 矩阵需要取 128 列,它们在 K 方向上都会被切成若干个长度为 KK 的小段。KK 一般取 16 或 32,这样 A 的一个小片是 128x16,B 的一个小片是 16x128,加载到共享内存的开销可控。
对于线程块内部,每个线程还要再负责一个 TM x TN 的微块。比如每个线程计算 8x8 的累加块,那么一个 128x128 的输出块就需要 16*16 = 256 个线程。这是非常常见的划分方式。分块本身不复杂,但它决定了后面所有访存与流水线优化的边界条件。
2.2 双轨计算路径:Tensor Core 与普通 FMA 各司其职
现在的加速卡普遍具备专用矩阵乘指令,通常称为 Tensor Core。这类指令可以用一条指令完成一个小型矩阵乘加运算,比如常见的形态是 16x8x16 或 16x16x16,输入精度常见为 FP16、BF16、TF32。相比标量乘加指令,Tensor Core 能在一个周期内完成多得多的计算量,是提升峰值吞吐的核心手段。
DeepGEMM 的主路径就建立在 Tensor Core 指令之上。以 FP16 输入、FP32 累加为例,指令要求参与计算的矩阵片段在寄存器里有特定的排布方式,有的要求按列分片,有的要求按行分片,这个细节直接影响从共享内存加载数据时的索引顺序。我会专门设计一个“数据 swizzle”阶段,把从全局内存拿到的普通排布数据,在写入共享内存时变成指令友好的布局。
但 Tensor Core 并不是万能的。比如某些自定义精度格式,或者当矩阵边界不是指令形状的整数倍时,Tensor Core 指令根本没法用。所以我保留了一条普通 FMA 路径,专门处理尾块和不支持的精度组合。这条路径不求速度有多快,只求正确性没有问题,反正尾块占整体计算量的比例很小。
2.3 共享内存布局与 Bank Conflict:一个被低估的性能杀手
写 GEMM 内核时,共享内存的访问模式几乎决定了一半的性能。共享内存被硬件组织成 32 个存储体(bank),每个 bank 的宽度通常是 4 字节。硬件保证:如果同一个 warp 里的 32 个线程访问的地址恰好分别落在 32 个不同 bank 上,可以一次完成;但如果两个或更多线程访问同一个 bank,硬件就要分多次处理,这就是 bank conflict。
举例来说,假设线程 t 想读取共享内存地址 base + t * 4(每个线程读 4 字节),那么地址依次落在 bank 0、1、2...31,完全无冲突。可如果地址是 base + t * 8,那么线程 0 读取 bank 0,线程 1 读取 bank 2,线程 2 读取 bank 4,一次下来有一半的 bank 被踩中两次,实际的访存吞吐直接腰斩。
在矩阵乘里,A 分块在共享内存中经常按“行优先”方式存放。当不同线程读取同一行中不同列的数据时,连续列会映射到连续 bank,看起来还行;但读取下一行同一位置的数据时,如果行的字节数恰好是 bank 数的整数倍,那么每一行同一列的地址会落在同一个 bank 上,变成灾难级的冲突。解决手段主要有两种:一是给每行加 padding,让行的实际宽度不等于 32 的整数倍;二是做 swizzle,比如把地址的低位做异或重排,让同一 warp 的 32 个线程尽量分散到不同 bank。DeepGEMM 里两种方法都用了:A 分块主要用 padding,B 分块因为访问 pattern 更复杂,用了一个简单的异或 swizzle,实测能把共享内存访存时的 bank 冲突降到几乎为零。
这里有个常见误解:很多人以为共享内存容量大就万事大吉,忽略访存宽度。实际上共享内存的带宽远高于全局内存,但依然有限。一个 warp 如果发生 2 路冲突,共享内存那一步的耗时就会翻倍,进而拖慢整个流水线。所以我建议任何写 GEMM 内核的人,第一件事就是在性能分析工具里看 shared memory bank conflict 的统计,这个数字会告诉你很多问题。
3. 让计算核心“吃饱饭”:流水线调度与尾块处理
3.1 从串行加载到三阶段流水线:双缓冲为什么是必需品
一开始我写了一个逻辑正确但很慢的版本。流程是:先把 A、B 的一个分片从全局内存加载到共享内存,然后等所有线程都同步好,再读回寄存器做矩阵乘,等计算完,再加载下一片。这个串行结构下,加载数据时的长延迟会完全暴露,计算单元经常处于等待状态。
后来改成三阶段流水线。三个阶段分别是:加载阶段、计算阶段、写回阶段。当计算当前分片时,异步地把下一个分片从全局内存搬到共享内存的另一块缓冲区;等当前分片算完,两块缓冲区角色互换。这就是典型的双缓冲,也叫软件流水线。
现代加速卡通常提供异步拷贝指令,能够在不占用寄存器的情况下,把全局内存数据直接传到共享内存。调用者发出拷贝请求后,不需要阻塞等待,而是继续执行后续指令,之后在要用数据时再统一等待全部完成。这个机制配合双缓冲使用效果非常好:K 循环每往前走一步,计算单元都在连续工作,访存和计算完全重叠。
用伪代码表示核心循环大概是这样的逻辑:
// 以下为思路示意,不含具体平台指令 for (int k0 = 0; k0 < K; k0 += KK) { // 预取下一次要用的数据到备用缓冲区 async_copy(A_smem[1 - cur], A_global + tile_offset_a, bytes_a); async_copy(B_smem[1 - cur], B_global + tile_offset_b, bytes_b); // 等待当前缓冲区数据就绪 wait_all_async_copies(); // 每个线程从当前缓冲区读回寄存器并执行矩阵乘 compute_tile(A_smem[cur], B_smem[cur], accum); // 交换缓冲区 cur ^= 1; }实际实现中还有不少细节,比如异步拷贝请求的提交与等待组管理、在循环最后阶段避免重复预取等。整体上,双缓冲是 GEMM 内核从“能跑”到“跑得快”的第一个关键门槛。
3.2 尾块分支:不让边界条件拖累主路径
当 M 或 N 不是 tile 大小的整数倍时,总会剩下一圈不完整的输出块。如果在主循环里加一个 if 判断“是否有越界”,虽然逻辑上正确,但会污染整个主循环的指令流,让所有线程在每个分片上都多一次分支判断,性能受损。
我的做法是把尾块彻底分离出去。启动两个核函数:一个专门处理 M 和 N 都是 tile 整数倍的主区域,核函数内部没有任何边界检查;另一个专门处理边界不齐的尾块区域,使用更小的 tile 尺寸,甚至可以走普通 FMA 路径。实际测试下来,把尾块分离之后,主循环的吞吐提升了差不多 5%,这点收益在超大矩阵上非常可观。
尾块本身的计算量不大,但要注意越界加载的问题。当 K 维度也存在尾块时,同样要单独处理:不能假装 K 被填充成整数倍,否则会读到错误数据。我在初始化共享内存时会对越界区域填零,然后让尾块核函数执行合法范围内的计算,确保结果正确。这里也曾付出过不少调试时间,后面会细说。
3.3 流水线级数与寄存器压力:不是缓冲越多越好
流水线级数是一个需要权衡的参数。三阶段流水线可以隐藏一部分延迟,但如果计算时间比加载时间长很多,可能两个缓冲就够了;反过来,如果加载时间比计算时间长,可能需要更多缓冲才能平滑掉访存波动。理论上,流水线级数越多,越能容忍访存延迟,但每多一级就要多占一份共享内存和更多的同步逻辑。
在 DeepGEMM 的不同阶段,我试过 2 缓冲、3 缓冲、4 缓冲。在目标硬件上,4 缓冲的效果最好;换到另一代硬件上,3 缓冲反而更优。原因在于不同硬件各类指令的延迟和带宽比例不同。这个结论没办法从书本上直接查到,只能靠实测。一般建议是从 2 缓冲开始,逐步增加,同时观察设备利用率(通过性能分析工具看计算单元的忙闲比例)。
共享内存总量是有限的。当 BM = 128、BN = 128、KK = 32 时,A 分片需要 128322 字节(FP16),B 分片需要 321282 字节,两个分片占共享内存 16KB。如果做 4 缓冲就要 64KB,占一块芯片共享内存的相当大比例,此时留给其他数据结构的空间就很少了。这也是为什么需要平衡 tile 大小和缓冲级数。
3.4 一个真实的血泪教训:第一版为什么慢得离谱
我第一次把逻辑正确的内核放到性能分析工具里看,设备利用率只有不到 20%,吞吐只有理论峰值的 7%。当时的第一反应是“指令太复杂”,但实际上问题出在访存上:共享内存存在大量 2 路甚至 4 路 bank 冲突,异步拷贝完全没有使用,每次加载都让计算单元停下来等待。后来逐步修复了这三个问题,性能从 7% 一路涨到 40% 多。整个过程中最有效的不是一上来就去抠 Tensor Core 指令,而是确保数据搬运路径是顺畅的。想让计算核心“吃饱”,首先要让人家能一直有活干,而不是干一秒等三秒。
4. 实测数据与调优过程:从“比库慢三倍”到追平再到反超
4.1 第一版的数据与问题定位
第一版 DeepGEMM 在一个 4096x4096x4096 的 FP16 矩阵乘上跑出了 58 TFLOPS。作为对比,同一硬件上厂商通用库能跑到 620 TFLOPS 左右。这个差距已经不是“慢一点”而是“差一个数量级”了。我分析了性能分析工具给出的指标:
- 共享内存 bank conflict 率达到 37%,每个 warp 在访问共享内存时平均要多花一半等待周期;
- 没有使用异步拷贝,每次加载数据后必须阻塞等待,计算单元空闲率高达 78%;
- 寄存器分配不合理,有溢出到局部内存的情况,等于把寄存器数据又存了一遍。
这三点里,第三点尤其容易被忽略。当寄存器压力过大时,编译器会把多余变量自动放到局部内存,而局部内存在物理上还是走全局内存的路径,性能损失极其明显。检查方法很简单:看编译生成的寄存器用量,如果超过可用寄存器数量,就需要缩小每个线程的微块尺寸或者简化指令逻辑。
4.2 迭代优化路径与每一步的收益
后面的调优是一个版本一个版本叠加出来的。我整理了一张表,记录主要改动的效果:
| 版本 | 关键改动 | 实测吞吐(TFLOPS) | 占理论峰值比例 |
|---|---|---|---|
| v0 | 正确性验证版 | 58 | 约 7% |
| v1 | 共享内存 padding + swizzle | 121 | 约 15% |
| v2 | 异步拷贝 + 双缓冲 | 342 | 约 42% |
| v3 | 调整 tile 大小与线程微块(BM=128, BN=128, KK=32,每线程 8x8) | 490 | 约 61% |
| v4 | 多级流水线(4 缓冲) + 分离尾块 | 652 | 约 81% |
| v5 | 针对固定形状做关键参数搜索(JIT 选择最优组合) | 724 | 约 90% |
这个表能说明一个很核心的观点:GEMM 性能不是靠某一个“绝招”拉上去的,而是每一层都扣掉一点浪费之后,累积出来的结果。单独看每一步,好像提升都不算惊世骇俗,但叠加起来就是从 7% 到 90% 的巨大差距。
v3 里调整 tile 大小这个改动给我印象很深。原来每个线程负责 16x16 的微块,寄存器用量过大,编译器开始溢出。改成 8x8 之后,寄存器占用显著下降,局部内存访问消失,同时每个线程块还能容纳更多线程,让计算密度更均匀。tile 大小不是一个可以“照着论文抄”的参数,它跟寄存器数量和共享内存容量强相关,需要在实际设备上搜索。
4.3 与厂商通用库的最终对比
最终版本和厂商通用库在同一硬件上的对比结果如下(FP16 输入、FP32 累加、M = N = K = 4096):
- 厂商通用库:约 620 TFLOPS,约为理论峰值的 77%;
- DeepGEMM:约 724 TFLOPS,约为理论峰值的 90%。
坦白说,每代硬件的上限和厂商库的优化程度都在变,今天能领先,换一个形状、换一个版本可能就没有优势了。但对一个固定场景来说,16% 的吞吐提升意味着同样的训练或推理任务可以节省约 14% 的时间,这个收益是实打实的。
另一个得到验证的好处是融合算子。DeepGEMM 在矩阵乘内部加了一个简单的 bias 和 scale 参数,结果写回之前直接完成缩放。把之前需要三个 kernel 完成的操作压缩成一个,端到端反而减少了约 20% 的耗时,这部分收益比单纯矩阵乘提速更明显。
4.4 用性能分析工具找到的一个隐藏热点:L2 命中率
有一轮性能始终卡在 65% 附近,计算单元占用率不低,但总吞吐上不去。后来看 L2 缓存命中率,发现只有 78%,对于矩阵乘这种访存模式相对规则的操作来说太低了。原因是我在遍历输出 tile 时,A 矩阵的 tile 顺序和 B 矩阵的 tile 顺序没有对齐,导致相邻线程块访问的全局数据在 L2 里不重合。
调整 grid 的遍历顺序,让同一个 k 切片上相邻的两个输出块尽量同时调度,L2 命中率回升到 90% 以上,最终吞吐也突破了 80% 大关。这个经验可能很多人不知道:GEMM 内核不仅要关心共享内存和寄存器,也需要关心 L2 缓存的行为。尤其在现代加速卡里,L2 的容量和带宽比想象中更重要。
5. DeepGEMM 能走多远:边界条件与长期规划
5.1 当前支持的能力边界
目前 DeepGEMM 支持的精度包括 FP16、BF16 和 TF32,累加统一用 FP32。支持自定义的 bias 与缩放因子,方便做混合精度训练中的损失缩放。矩阵形状方面,针对 M = 1024 到 8192、N = 512 到 4096、K = 512 到 4096 之间做了重点调优,超出这个范围的形状也能跑,但性能不会是最优。
这里要说明,BF16 和 FP16 的指令路径并不完全相同。FP16 在 Tensor Core 指令里有更成熟的 16x8x16 形态,BF16 有时要走另一种格式组合,共享内存里的排列方式要相应调整。代码里我用一个编译期常量控制这两个路径,避免运行时分支造成额外开销。
5.2 可移植性与 JIT 化的拉扯
每次把 DeepGEMM 搬到不同架构的硬件上,都要重新测一遍最佳参数组合。这个“重新测”的过程最开始需要手工改代码、重新编译,非常浪费时间。后来我决定把 BM、BN、KK、流水线级数、线程微块大小全部变成配置项,在第一次运行时通过一段简单的搜索代码,用几百次小规模运行选出最优组合。
这会引入一段 JIT 编译的过程,但收益非常大:同一份源码,在不同代际硬件上都能自动适应,而不需要我维护多套手工调优分支。这个方向也符合现代算子库的发展趋势——与其在源码里写死参数,不如让计算框架在启动时根据硬件特征自动生成最适合的内核配置。
不过说起来容易,做起来还需要处理几个麻烦点:搜索空间的控制(比如只搜索有限几组典型参数,而不是全排列)和搜索时间预算(不能为选参数花掉比节省下来的运行时间还久)。DeepGEMM 里的做法是先把明显不合理的组合过滤掉,再用 5 组左右候选做实际计时比较,通常能在几次启动内收敛。
5.3 给同样想自研 GEMM 内核的朋友几条建议
如果看完这篇你也想动手写一个自己的 GEMM,我有几条基于血泪的经验想分享:
- 先追求正确,再追求快。第一版哪怕只有 7% 的峰值利用率都没关系,先把结果算对,建立完整的测试用例,尤其是边界形状和异常参数,后面优化才不会为错误找半天借口。
- 不要在开头就直接写底层的微码或汇编。标准的高层编程模型足够好,先把分块、流水线、swizzle 这些思路验证清楚,再去考虑极个别指令级优化。
- 一定要学会读性能分析工具。bank conflict、寄存器溢出、异步拷贝等待周期这些指标,比任何经验都更直接。拿到数据再改代码,而不是靠猜。
- 固定形状先做深。不要一开始就想支持任意 M、N、K,那样调度逻辑会复杂到你失去耐心。
- 保存好不同版本的性能记录。每一次改动带来了什么,为什么带来自,至少要有一个简单表格,否则优化到后面会忘记哪个技巧真正有效。
我目前的工作流里,DeepGEMM 还处于“为固定模型服务的私有算子库”阶段,后续计划把注意力里的 QK^T、PV、残差连接、LayerNorm 也逐步融合进去,做成一个更完整的推理算子集。不过这一步每往前走一点,都需要先确认计算图层面能拿到足够的收益,否则就是在给代码库增加复杂度。
写这个项目的过程中我最大的体会是:真正难的不是理解分块矩阵乘法的数学,也不是记住几条硬件指令,而是在成千上万个参数组合里找到那条最平滑的路。每当你觉得“应该够快了吧”,性能分析工具总能再给你指出一个偷偷等着拖延你的隐藏瓶颈。如果你也在做类似的事情,希望这篇分享能让你少走几个我已经替大家踩过的坑,哪怕只省下一次调试的时间,也算值得。