DeepGEMM内核优化实战:从分块、流水线到硬件极限
今年我大部分时间都耗在“把矩阵乘法跑得更快”这件事上。项目代号就叫“DeepGEMM”——不是某个大团队的产品纯粹是我在业余时间写的一个极致优化版 GEMM 内核。GEMM通用矩阵乘法听起来很底层、很老古董但现代深度学习训练和推理里有大约 70% 到 90% 的计算量最终都落到矩阵乘法上。无论是全连接层、卷积通过变换转成 GEMM还是注意力机制里的 QK^T 和 PV归根到底都是在做“一批矩阵乘”。因此能不能把这一件事做到接近硬件极限直接决定了上层模型的性价比。说句实话市面上已经有很多成熟的 GEMM 库性能也调得相当好。那为什么还要自己写一套因为“看懂别人代码”和“自己能在编译器和硬件之间讨价还价”是两码事。DeepGEMM 的目标不是“赢过某个大厂闭源库”而是搞清楚一块显卡从理论峰值到实际可跑性能之间到底被哪些因素截胡了。这篇文章就把我调内核的真实过程、踩过的坑、以及最后沉淀下来的一套可复用优化思路完整拆出来适合正在学高性能计算、写算子内核、或者打算在深度学习框架里做自定义算子优化的朋友。1. 为什么把矩阵乘法单独拎出来“Deep”一把1.1 GEMM 是深度学习的算力底座先给不常写底层的朋友补个背景。GEMM 的数学定义特别简单C alpha * A * B beta * C。其中 A 是 M 行 K 列B 是 K 行 N 列C 是 M 行 N 列。三个循环就能做出一个“能用但不快”的版本for (int i 0; i M; i) for (int j 0; j N; j) for (int k 0; k K; k) C[i][j] A[i][k] * B[k][j];这一段代码放到现代 GPU 上跑性能大概率只有硬件峰值的 2% 到 5%。原因很多但最致命的一条是每次都去读全局内存相当于每天几千次跑去仓库取货而且同一个数据会被翻来覆去地读。想象一下楼上超市每天早上要从仓库搬进来 1 万箱矿泉水如果每卖给一位顾客就专门跑一趟仓库取一瓶物流就崩了。GEMM 优化的核心本质上就是“如何把仓库的货尽量多、尽量整齐地搬到楼下货架”并让货架上的货在卖掉之前不要反复往仓库跑。媒体上天天说的“某芯片的 AI 算力是多少 TFLOPS”指的就是硬件在理想条件下每秒钟能完成的浮点运算次数。但理论峰值只有在数据排列、指令发射、访存模式全部“完美”时才有可能碰到。GEMM 计算强度高、访存模式规整属于最容易贴近峰值的算子之一。如果连 GEMM 都跑不到理论值的一半那其他更复杂的算子就更不用想了。这也是大家把 GEMM 单独拿出来深度优化的根本原因它是指标标尺也是“照妖镜”。1.2 硬件演进从乘加单元到张量核心过去十年主流 GPU 的变化不仅仅是“浮点运算器件变多了”更重要的是“运算粒度变大了”。最早做 GEMM 是在 CUDA Core 上一次算一个乘加FMA一条指令对一个线程内的两个数做运算。现在的主流架构里出现了专门做矩阵运算的硬件单元——张量核心Tensor Core——一条指令直接把一个 16×16 的 A 子块和 16×8 的 B 子块乘起来累加。指令粒度变大之后优化逻辑也跟着变了不再是把“线程”当成核心调度单位而是把“子矩阵块”当成第一调度单位。这是 DeepGEMM 项目最有意思的地方同样一个数学问题因为硬件调度粒度不同写出来的代码结构完全不同。老式写法是“每个线程搬一个数算一个数”新式写法是“每个线程负责搬运一块数据再由张量核心成批计算”。后者需要更精细的数据布局、更深的异步流水线以及更小心地躲开存储体冲突。这篇文章后面讲到的多数优化换到没有张量核心的老硬件上会失去意义。但理解了新式写法再回头看我之前为老硬件写的内核会非常清楚地理解“为什么那时候只能跑那样一个性能”。2. 从算法到硬件的三层映射GEMM 优化坐标2.1 Roofline 模型先算清楚天花板在哪很多时候调优失败不是因为技巧不够而是根本不知道天花板在哪里。Roofline 模型给了最简单的估算框架。对于 GEMM 这种算子我们关心两个指标计算强度每读一个字节数据能换来多少次浮点运算和访存带宽。GEMM 的计算强度可以这样算一次 MNK4096 的方阵乘法总计算量约 2 * 4096^3 次 FLOPs总共需要从全局内存把 A、B、C 读进来再写出去数据量约 3 * 4096^2 * 4 字节按 FP32 算。算下来计算强度大约是 682 FLOPs/Byte。把这个数画到 Roofline 图上它远在“带宽受限”区域之外所以 GEMM 本质上是“算力受限”的算子优化目标应该聚焦在提高计算单元利用率而不是压缩访存量。但“理论上不受访存限制”不等于“随便写就能行”。现代 GPU 有非常复杂的缓存层次全局显存相当于大仓库→ 二级缓存相当于区域分拨中心→ 共享内存相当于超市楼下货架→ 寄存器相当于收银员手边的零钱柜。每一层带宽逐级上升容量逐级下降。GEMM 优化的核心坐标就一句话让数据在不同层次之间以“整块、可预测、可流水”的方式流动并且尽可能在高层寄存器、共享内存完成最多次数的复用。2.2 三层并行的直观拆分GEMM 在硬件上天生适合做三层并行Block 级CTA 级把一个大的输出 C 矩阵切成多个小块比如每个线程块算 128×128 的 C 子块。不同线程块之间完全独立互不通信可以并行分发给不同的计算单元。Warp 级线程块内部的多个 warp 再对 128×128 的子块做进一步切分。比如一个 block 里塞 8 个 warp每个 warp 负责 64×32 的输出小块。指令级在 warp 内部直接用张量核心的矩阵指令完成小规模子矩阵乘加。这一层最关键因为指令面对的已经是硬件原生支持的 16×16×8 或 16×8×16 之类的小矩阵块。我最初写 DeepGEMM 时犯过的最大错误就是试图把“线程块切多小”和“warp 内部布局”分开考虑。后来发现硬件指令的矩阵形状会直接约束线程块切分方式。如果你选的 block tile 尺寸不能完整容纳指令需要的数据块就会被迫在寄存器之间做多余的 shuffle 和搬运性能立刻掉一个档次。三层映射必须“自顶向下逐层倒推”先确认张量核心指令的形状再决定 warp 的职责最后定义 block tile 的边长。用个不严谨的类比食堂备菜。全局内存是大农场共享内存是备菜台寄存器是炒锅。指令执行节奏就是炉火。一个熟练的师傅不会等客人点菜之后才去农场摘菜而是提前把菜备好放在备菜台同时手里这一锅还在炒。这就是后面要说的流水线pipeline。而“三层映射”则是在回答另一个问题到底哪一片菜归哪个师傅炒、多少人管一个灶、灶和灶之间怎么不打架。顺序反了后面所有优化都是白搭。3. 关键优化手段逐个拆解分块、流水线与存储冲突3.1 分块策略形状决定命运DeepGEMM 最终选定的 block tile 是 128×128每个线程块内部布置 8 个 warp每个 warp 负责 64×32 的输出片K 维每次推进 16或者 32看精度和指令形状。为什么是 128×128不是 256×64也不是 64×256答案要从“数据复用”和“寄存器压力”两个角度看。设 block tile 为 B_M × B_NK 维每次推进 B_K。那么一个线程块在计算当前 K 切片时需要把 B_M×B_K 的 A 子块和 B_K×B_N 的 B 子块搬进共享内存。中间结果 C 子块的大小是 B_M×B_N。如果 B_M×B_N 太大每个线程要维护的累加器寄存器数量就会爆掉。GPU 寄存器的总数是有限的每个线程最多 255 个寄存器左右是常规上限一旦超了就会产生“寄存器溢出”数据被“挤到”本地内存性能断崖式下跌。而 B_M×B_N 太小又会导致每次搬运 A/B 的开销占比过大吞吐上不去。128×128 是一个在大量主流尺寸上都能获得较高复用率和合理寄存器压力的“甜点值”。这里提到一个值得背下来的经验让 B_M、B_N、B_K 的尺寸和张量核心指令形状保持对齐关系。比如张量核心一次算 16×8×16你选的 warp 输出块 64×32 就能被指令切成 8 个 16×8 的子块K 维 16 也刚好和指令内 K 维长度一致。如果选了 66×34 这种“刚好差一点”的尺寸就会有一些寄存器宽度填不满或者边界处理时多出大量条件判断。高性能 GEMM 的尺寸设计没有魔法就是反复算“指令吞吐、寄存器占用、共享内存占用”这三个变量之间的平衡。3.2 寄存器级双缓冲与异步流水线先看经典的一版内核主体示意伪代码实际会比这个复杂得多float acc[TN][TN] {0}; for (int k 0; k K; k BK) { load_tile(A_block, B_block); // 同步把 A、B 子块搬进共享内存 for (int i 0; i TM; i) { for (int j 0; j TN; j) { for (int kk 0; kk BK; kk) { acc[i][j] A_block[i][kk] * B_block[kk][j]; } } } }这种同步式写法最大的问题load_tile执行期间计算单元在“干等”。现代 GPU 可以通过多线程块并行来掩盖一定程度的访存延迟但线程块内部的等待依然会拉低单块效率。解法是双缓冲或多缓冲预先把下一轮循环需要的 A、B 子块发一个异步拷贝指令例如cp.async当前这一轮照常计算。计算和搬运就像两条并行装配线等这一轮算完下一轮的数据已经就位了不需要再等。在我的实测里从同步加载改成 4 级流水线之后一个 4096×4096×4096 的 FP16 GEMM 在主流消费级 GPU 上从约 45% 峰值提升到了 73%这是一个性价比极高的改动。流水线级数stage 数也有讲究设太少掩盖不了延迟设太多共享内存装不下反而拖慢。一般 3 到 5 级比较常见。选 4 级属于折中方案一方面能覆盖大约 80% 到 90% 的搬运延迟另一方面共享内存占用也就比单缓冲多了三四倍通常还在预算内。“兵贵神速”这话在 GPU 上有点反直觉很多优化不是让指令跑得更快而是让指令之间的“等待”少一些。看内核性能先别数算力指令条数先看“每一周期里计算单元有百分之多少在干活”。3.3 共享内存的 Bank 冲突隐藏的性能刺客假如你的内核性能始终比预期低 20%而且怎么调参数都没有改善大概率是撞上了 Bank 冲突。共享内存是分“银行”的典型架构是 32 个 bank一个 warp 的线程同时访问共享内存时如果多个线程访问了同一个 bank 的不同地址硬件就要把请求拆成多次串行处理。表现就是同一段代码数据换个布局性能差一倍。传统方案里一个 warp 读 A 子块的某一列时线程的地址如果按列主序排布很容易让多个线程同时打到同一 bank。DeepGEMM 里我花了两个晚上排查这个问题最后用上了Swizzle 模式对共享内存地址做一个简单的异或变换让同一 warp 内的相邻线程访问分散到不同 bank。伪代码看起来像// 原始地址 int addr row * stride col; // 经过 swizzle 之后 int swizzled_addr (row * stride) ^ ((col 7) * 8);这里swizzled_addr里的异或会让列方向上的数据“错开”分布。代价是地址计算多了一两条整数运算和 Bank 冲突带来的性能损失相比完全值得。另外还要注意用cp.async异步搬运到共享内存时写入端也可能有 Bank 冲突。如果一个 block 里的线程按顺序连续写通常没问题但如果你人为做了复杂的数据重排要检查写入端的地址模式。读端和写端是两套逻辑最好都单独测一遍。3.4 数值精度从 FP16 到 TF32 的取舍账现代 GEMM 不是只调“速度”还要调“数值正确性”。DeepGEMM 里的核心循环我用了 FP16 做乘法、FP32 做累加。这样做的原因是FP32 乘法的硬件吞吐量远低于 FP16通常只有几分之一但纯 FP16 累加又容易在 K 维很长时产生明显舍入误差。FP32 累加相当于给每个“乘法后的小数”先换成更高精度再相加误差会小很多。对深度学习来说训练时一般也推荐用 FP16 乘法配 FP32 累加。另一类选择是 TF32 格式它本质上是截断的 FP32只有大约 19 位有效精度。TF32 的好处是能利用张量核心的高吞吐坏处是比普通 FP32 精度低一些。对于需要“全精度又想要速度”的场景TF32 是不错的折中但最终是否可用要靠误差测试说话。格式指数位尾数位典型用途注意事项FP16510深度学习训练/推理范围小容易溢出BF1687大模型训练范围同 FP32但精度低TF32810截断后混合精度训练近似 FP32FP32823通用计算/基准吞吐低于 FP16表格只是参考真正决定用哪种格式的是“你的算法能容忍多大误差”。DeepGEMM 里我保留了运行时格式切换的接口默认走 FP16 乘法 FP32 累加需要时也能切到 TF32。很多人一上来就把业务层的数据类型改了结果误差超标回头又怪库不靠谱。正确做法是底层算子支持多种输入类型业务层根据误差测试选。内核层只负责“在特定格式组合下靠近峰值”。4. 实测中避不开的坑性能回退与数值异常的排查链路4.1 性能神秘回退参数改裆性能反而崩了有一次我尝试把流水线 stage 数从 4 调到 6希望进一步掩盖延迟。跑起来之后发现性能不仅没涨反而从 73% 掉到了 61%。刚开始以为是修改引入 bug后来打开编译日志才发现问题出在共享内存占用。6 级流水线需要的共享内存超过了硬件的静态分配阈值编译器被迫把一部分数据放到“后门”内存访问速度差了一个数量级。这个坑在参数调优时非常典型一个方向性的优化加深流水线本身合理但没有同步考虑硬件资源上限。这个坑给我一个重要教训看内核的上限永远要同时画两条曲线一条是计算指令的理论吞吐一条是访存和资源约束下的可达到吞吐。两者中最低的那条才是真天花板。后来我每次调整参数都会先在脚本里打印一份寄存器数、共享内存占用、指令预估吞吐量的清单再跑基准。宁可慢一点也不要瞎猜。4.2 数值异常inf 和 NaN 的完整定位过程有一次 DeepGEMM 在单测里输出了 NaN但错误只在大矩阵尺寸下出现小矩阵完全正常这就很迷惑。我排查的顺序是这样的先检查 K 维不是 16 的整数倍时边界处理是否正确。大尺寸容易触发“边缘 tile 不满一个指令长度”的路径数组访问越界很可能读到垃圾值。我把if (k BK K)这类边界判断全部重写了一遍发现确实有一处是“先算后判”导致最后一轮循环里读到了未初始化的共享内存。再从浮点运算角度检查。K 很大时累加值数量级差异巨大FP16 中间结果可能溢出成 infinf 参与后续累加就成了 NaN。我把乘法结果强制转成 FP32 再累加情况有所缓解。最后查异步流水线的 wait 位置。cp.async发出去之后如果没等数据到位就开始算也会出现读未初始化数据。这个问题最难查因为有时候“碰巧”能对上时序大部分情况又都对偶尔才出错。解决办法是给流水线阶段加一个全局屏障grid-wide sync并在短尺寸下跑 500 组随机输入做压力测试。排查数值问题不要一个函数一个函数地盯先看边界再看精度最后看并发同步。按照这个顺序大部分“偶发异常”都能找到根因。4.3 排障工具与性能画像的正确打开方式很多初学者拿到 profiling 工具就只会看一个“平均 kernel 执行时间”。但内核优化的关键指标远不止这一个。我自己的习惯是先看四个维度实际算力利用率 实测 FLOPs / 理论峰值 FLOPs。别的都可以不看这个最直观。共享内存吞吐如果共享内存吞吐长期接近 80% 以上说明 Bank 冲突可能很严重或者访存模式有优化空间。L2 命中率GEMM 的理想情况是 L2 命中率尽量高。如果某个 shape 的命中率明显偏低说明 block tile 的调度顺序和局部性没配合好。频率曲线有些实测结果忽高忽低可能是供电限制或温度导致降频。先排除这个干扰项再谈代码优化。这四个指标都符合预期之后再去调指令级细节否则很容易浪费时间。比如你银行冲突已经让共享内存吞吐爆了却还在那里调整__syncwarp()的时机那就是跑偏了。用数据画像代替直觉判断是 DeepGEMM 项目里我做得最有价值的一件事。5. 从 DeepGEMM 到算子库下一步该怎么走5.1 经验迁移FlashAttention、卷积与稀疏场景GEMM 优化思路之所以值钱是因为它可以迁移到很多看似不同的算子。FlashAttention 就是把注意力机制的核心部分改造成带 Online Softmax 的 GEMMQK^T 的乘积计算是一个 GEMMPV 的计算是另一个 GEMM中间穿插按行缩放。我在 DeepGEMM 里写的小矩阵算子稍微改一改就能作为 FlashAttention 的内核组件复用。卷积则可以通过 im2col 变换补丁成一个大 GEMM虽然额外占内存但 GEMM 内核马上就能派上用场。稀疏化场景也类似结构化稀疏的权重矩阵可以切成密集小块每个小块依然走 GEMM 核心逻辑只是外层多了“跳过零块”的调度逻辑。这些迁移只有在你“亲自写过一遍底层 GEMM”之后才会变得顺手。没写过的人看 FlashAttention 源码像天书写过之后会发现它就是把“流水线”和“数值稳定”两个老朋友重新组合了一遍。5.2 自动调优思路让程序自己搜索参数空间手动调参终究有天花板。DeepGEMM 到了后期我开始把所有关键维度block tile 尺寸、warp 数量、流水线 stage 数、Swizzle 模式、是否用异步拷贝全部抽成模板参数然后用一套自动调优脚本去搜索。搜索策略也很简单先粗粒度随机采样一轮锁定几个候选区域再在候选区域做网格搜索。一轮下来通常能从手工调优的 73% 再推到 76% 到 78%。不要小看这几个百分点在数据中心里单算子核提高 5% 的利用率全年累计省下的电费和折抵的算力采购费用非常可观。还有一个已知但常被低估的结论没有一套参数能在所有矩阵尺寸上同时最优。方阵 4096 的最优解用在 M128、N65536 的瘦长矩阵上表现可能很差。所以算子库都会维护一张“shape - 参数”的映射表而不是硬编码一套配置。DeepGEMM 到这一步也从“一个内核”长成了“一个小型算子库”。5.3 工程化细节接口兼容与多设备适配写内核本身很有意思但如果目的是“给别人用”还差最后一步工程化。以 BLAS 风格的 GEMM 接口为例函数签名里那些transa、transb、alpha、beta参数一个都不能漏。很多人写了个快内核却只支持“自己约定好的方阵且 alpha1, beta0”的简化场景落到真实项目里立刻没法用。我第二版 DeepGEMM 就吃了这个亏有人拿它去算矩阵加法才发现beta参数根本没实现。后来我把这四个参数补全内核性能虽然下降了 1% 左右多了一些分支判断但“能用”比“好看”重要一百倍。多设备适配是另一个容易被忽略的坑。不同厂商、不同代际的硬件指令形状、共享内存大小、带宽比例都不一样。一套“最优”配置拿到别的设备上可能很平庸。好在核心优化逻辑是通用的——分块、流水线、Swizzle——只需要把参数重新搜索一遍代码改动不大。这也是我坚持把优化点全部封装成模板参数的原因。结尾一点个人体会最后说一点个人感受。DeepGEMM 写到现在最大的收获不是“我的内核比某个开源库快多少”而是我终于建立了“算法复杂度”和“硬件资源预算”之间的直觉连接。过去看到一个算子的计算公式我只能大致估它是内存密集还是计算密集现在我会下意识地去想它需要多大的寄存器文件共享内存能不能装下中间结果指令级并行度够不够掩盖访存延迟如果你也想动手做一个类似的“Deep”系列项目我的建议是先把环境跑通下一个最简单的三项循环版本测出它的性能再用 profiling 工具算出理论峰值然后一项一项地试分块、双缓冲、Swizzle。每完成一步回去看一下性能曲线把每一步背后的“为什么”写进注释里。这个过程比看一百篇优化文档都管用。关于后续方向我目前正在把 DeepGEMM 里积累的分块和流水线逻辑往更复杂的稀疏矩阵算子迁移等跑出稳定版本再来写一篇实战记录到时候见。