
做底层算子优化的人大概都绕不开这样一个念头明明调一个库就能拿到近乎峰值性能的矩阵乘法为什么还要自己去写一个“DeepGEMM”我最早接手这个项目时也是这么想的直到在某个推理引擎里连续碰到几个通用库搞不定的场景——batch只有1、K特别大、还要把激活融合进来我才真正下定决心自己从零搓一个GEMM内核。DeepGEMM不是一个复杂到看不懂的项目恰恰相反它的核心思想非常朴素围绕特定硬件和特定形状把通用矩阵乘法GEMM的性能榨到极限。这篇文章不打算讲那些花哨的推导而是把我从零实现到逐步调优的整套思路拆开包括分块策略、共享内存布局、指令调度、数值精度和真实调参记录。适合正在做算子开发、推理引擎部署、模型性能优化的人参考哪怕你之前只写过naive的矩阵乘法跟着这套思路也能写出可用的高性能版本。1. 为什么还要再写一个GEMM放弃现成库的理由1.1 库调用掩盖的三个真相很多人觉得高性能矩阵乘法直接调库不就行了自己写大概率写不过。这话一半对一半不对。通用库确实在“通用”这个尺度上做了极致的优化但它优化的是统计意义上最常见的形状、最普适的内存布局、最保守的数值策略。而你一旦进入真实业务场景这些“优化前提”可能全都不成立。第一个被掩盖的真相是通用库对形状并不敏感。它内部会做启发式选择碰到M1、N4096、K4096这样的推理场景大概率会走一条为“大M大N大K训练形状”设计的路径结果就是明明只有一两个block在干活整个设备空闲了一大半。第二个真相是库接口封装了内存布局但它假设你的数据是连续标准布局。一旦你的A矩阵来自某个量化算子输出、B矩阵是转置后拼接的权重你用库时往往得先做一次内存重排这个拷贝开销可能直接把计算带来的收益吃掉。第三个真相是库不给你融合机会。激活函数、缩放因子、偏置、Clamp这些在推理里基本必然存在但放在通用库里就成了多次读写的开销。我做DeepGEMM的动机说白了就三个字要可控。我要能在内层循环里顺手把LeakyReLU做了能在K维分块时动态加缩放能针对长K小batch的形状单独走一条快速路径。这些需求不是通用库能优雅满足的。1.2 通用最优不等于场景最优这里有个很容易被忽略的点同样一个GEMM“通用库最优”和“你的场景最优”经常不是同一种实现。直接调库跑起来可能已经很快了比如达到了理论峰值的70%你觉得自己没必要再折腾。但注意通用库为了平衡所有形状会在内核选择上做大量分支它在你的特定形状上可能只有60%或65%的效率而那丢失的百分之几到十几在端到端推理里往往就是批处理吞吐量的差距。我在DeepGEMM里做的第一个决策就是放弃“什么形状都最优”的幻想明确写死几类目标场景训练场景下M通常较大比如512到4096推理场景下M可能只有1到64但N和K往往很大。这两类场景我会走不同的主循环和分块参数而不是硬套同一个kernel。这个决策一开始看起来很“笨”但实测下来它带来的收益比后面任何一个指令级优化都大。正是因为放弃了通用性我才能针对M1的场景把整个C矩阵切成一条长条让每个线程块沿着K方向顺序扫避免在M维度上产生无效block和同步开销。而通用库不敢这么做因为它要先判断M1是否值得专门优化判断逻辑本身也会引入开销。1.3 DeepGEMM的定位平衡点在哪所以DeepGEMM并不是要取代现成的高性能库它更像是一把手术刀针对特定形状和融合需求做定制。我在项目里保留了naive版本作为正确性基线也保留了面向大训练形状的“通用快速路径”再叠加一个面向小batch长K的“推理专用路径”。这三层代码共用同一套分块框架但关键参数和循环结构不同。这样的结构还有一个隐性好处调试定位变得容易。当性能或精度出问题时我不用在几千行汇编式内核里大海捞针而是先判断当前形状走的是哪条路径再缩小到对应路径的逻辑。如果你也想动手写自己的DeepGEMM建议从一开始就保持这种“快速路径通用路径”的分层别急着把所有东西揉进一个kernel。2. 从计算强度反推设计先算账再写代码2.1 别问快不快先问瓶颈在哪开始写分块代码之前我先做了一道算术题。GEMM一次乘加需要2×M×N×K次浮点操作而需要搬进计算单元的数据量是M×N N×K M×KC矩阵读改写按分块局部化后可以压缩这里先按全量估算。两者一比就得到计算强度计算强度 ≈ 2MNK / (MN NK MK)这时候结论非常明显当M、N、K都很大时计算强度很高一个矩阵乘法本质上是计算密集型的性能的天花板由FMA吞吐决定而当M很小比如推理场景的M1或者K特别小的时候访存开销开始占主导你优化指令流水线可能远不如优化内存搬运来得有效。这个账必须先算清楚因为它直接决定了DeepGEMM的主循环怎么写。对于大矩阵核心任务是把数据喂给计算单元让FFMA永不空转对于小矩阵核心任务变成减少读入的数据次数和让C矩阵尽量留在寄存器或共享内存里。两种形态的代码长得完全不一样。2.2 三层分块寄存器、共享内存、L2的接力赛一旦确定瓶颈是计算接下来的问题就是怎么把M×N×K的大循环拆成能塞进硬件结构的小块。我的做法是三层分块每一层对应一级存储。第一层是Block Tile。整个C矩阵按128×128的块划分每个线程块负责一个块这个尺寸是给L2缓存和线程块并行度用的。第二层是Warp Tile。一个128×128的Block Tile交给4个warp处理每个warp拿32×32的子块。第三层是Thread Tile。每个warp内部有若干线程把32×32继续切成每个线程负责的微块我常用的配置是每个线程用8×8大小的累加器矩阵。为什么非要三层因为每一级存储的容量、带宽、延迟都不一样。寄存器最快但容量最小共享内存次之L2再次。三层分块本质上是让数据在这三级存储里逐级接力——A和B的大块从全局内存搬到共享内存再由共享内存搬进寄存器而C的累加结果尽量在寄存器里累积完最后再一次性写回。这也是很多高性能GEMM与naive版本最本质的区别naive版本每个线程直接去全局内存读A和B每算一次乘加就要发起一次访存而高性能版本是把数据先“囤积”在靠近计算单元的存储里一次搬运、多次复用。2.3 块形状和尺寸的试错起点很多第一次写GEMM的人会到处抄别人的Tile参数抄完发现性能稀烂然后怀疑是自己机器不行。其实尺寸这个东西非常依赖具体硬件的手写特性别人调好的参数未必适合你的场景和硬件代次。我的建议是找一组合理的起点然后做参数扫描不要凭感觉一锤定音。我最初用的一组参数是这样层级尺寸说明Block Tile (BM×BN)128×128照顾L2容量和块间并行度K维单次迭代 (BK)8或16决定共享内存单次搬运量Warp Tile (WM×WN)32×32每个warp负责的C面积Thread Tile (TM×TN)8×8每线程持有的累加器个数Warp个数4一个Block内4个warp这个起点不是拍脑袋而是基于硬件资源算出来的每个线程8×8需要64个累加寄存器加上A、B的片段寄存器、地址计算、临时变量寄存器用量在150到200之间不会爆掉寄存器文件又有足够的ILP指令级并行隐藏延迟。后续扫描我一般从“BM/BN是否变大”“BK是否太小导致搬运频繁”两个方向展开这个我放到后面调优章节细说。3. 内核循环的主体结构K维主序与微内核展开3.1 主循环里每个Stage在做什么DeepGEMM的主循环是沿着K维推进的不是沿着M或N维。原因很简单C矩阵的每一个点都是A的某一行和B的某一列的K维点积沿K方向一次推进就能同时更新当前Block Tile内所有C元素这个特性让K维主序天然适合做数据流水。主循环里每个Stage做四件事第一步从全局内存加载A和B各自的一段到共享内存记为阶段P0第二步等待共享内存就绪这叫同步点第三步每个线程从共享内存读入自己需要的A片段和B片段到寄存器阶段P1最后执行一批FMA运算把这组数据乘累加到自己的Thread Tile上阶段P2。整个结构用伪代码表达大致是P0: load A_tile[by][kk] - smem_A load B_tile[kk][bx] - smem_B __syncthreads() P1: load sA[ty][tx] - reg_A load sB[ty][tx] - reg_B // 按tile坐标取数存在swizzle时需做索引映射 P2: for i in 0..TM-1: for j in 0..TN-1: acc[i][j] reg_A[i] * reg_B[j] 前进kk回到P0这个循环写好之后我第一时间就不是看性能而是先拿小矩阵做数值对比。因为GEMM一旦分块最容易出错的就是索引坐标映射一个分块内部的对齐错位可能只在特定形状下触发非常阴险。3.2 双缓冲与组播把等待变成计算P0到P2的串行结构有个明显问题加载共享内存是慢的而计算是快的如果每次都要等共享内存加载完才开始算计算单元就会频繁空转。解决办法是双缓冲。我在共享内存里给A和B各准备两块buffer当前正在计算的那块记为当前buffer下一轮要用的数据提前加载到另一块buffer。这样P2的计算可以和下一次P0的加载重叠GPU在做乘加的同时DMA单元已经在下一次数据了。同步点也从一个变成两个一个确保上一轮计算结束前不能覆盖当前bufferconsumer同步一个确保下次数据写完后才能开始下一轮计算producer同步。组播multicast或者说广播读取是另一个关键。同一个warp里的线程在K维推进时A的某一行会被多个线程重复读取B的某一列也是。与其让每个线程都去共享内存读一遍不如让硬件感知到这是同一条数据一次读出来广播给多个线程。这个优化不在代码层面显式出现但你的数据布局和索引模式会影响硬件能不能自动识别这种广播模式。这也是为什么我坚持用规整的、线性偏移的索引方式而不是随手做各种奇怪的重排。3.3 微内核中的指令级流水主循环里最内层的P2部分我把它称为微内核micro-kernel。这一小段代码写得好不好直接决定了最终性能能到峰值的百分之多少。以每线程8×8累加器为例微内核要做64个乘累加对应64条FFMA指令。这里有个经验FFMA指令虽然是一个周期就能发出一条但如果编译器生成的指令流是“反复往同一个寄存器上累加”会产生很长的依赖链——下一周期要用的值必须等上一周期的结果算完流水线就堵住了。所以我写微内核时会让连续的几次FFMA依次作用在8个不同的累加器上把一个64步的长依赖链打散成8条独立的短链。这个思路有点像流水线上一个工人只拧一颗螺丝会卡住整条产线不如让几个工人各管一条螺钉线。实际我通常会让寄存器里的A片段和B片段交错排列先算acc[0][0]、acc[0][1]、acc[0][2]再算acc[1][0]……绕一圈回来再轮到acc[0][0]的下一次累加。循环展开因子选4或8太小依赖链压不住太大会把I-Cache挤爆具体数值要看指令发射宽度。4. 访存优化的硬仗共享内存冲突与Swizzle4.1 一个反直觉的停顿我的DeepGEMM在第一版完成时性能只比naive快了一倍离预期差得远。我拿着profiling结果一行行找发现计算单元利用率不低但是共享内存的吞吐指标异常高像是有东西在共享内存上来回打转。后来才意识到是bank conflict在作祟。共享内存按bank组织同一bank同一周期只能服务一次访问如果同一warp的多个线程同时访问同一个bank硬件就得把这些访问串行化。第一个版本我图省事A矩阵按行主序直接铺在共享内存里结果一个warp里的线程去读相邻列的时候恰好踩到了相同的bank或bank组访存效率几乎腰斩。这个问题的反直觉之处在于从高级语言看每次线程访问的都是“自己的地址”编译器也没有任何报错但性能就是上不去。如果不熟悉bank机制很可能反过来怀疑是自己循环展开写坏了白白浪费好几天排查时间。4.2 XOR Swizzle解决冲突的直觉解决bank conflict最通用的一招是swizzle。我的做法是对共享内存里的数据做一次XOR映射让原本会连续访问的地址被打散到不同的bank上。直觉上可以这样理解共享内存有32个bank同一warp有32个线程如果让地址的排列方式满足“每个线程访问的bank号互不相同”冲突就消失了。对于A矩阵的共享内存布局我会在索引里加一个XOR操作把行号和列号的一部分做一个异或再决定实际地址偏移。这个操作是在把全局内存数据写入共享内存时就做好的计算阶段读取时直接用映射后的地址几乎不增加额外指令开销。当然swizzle不是万能药。如果你的访问模式本身就比较随机XOR反而会引入新的碰头概率。我用它主要是因为GEMM里A片段的访问模式非常固定——同一个warp读同一行或同一块列规律性强swizzle的效果是立竿见影的。关键是要保证映射是双射也就是不丢失数据这个我会先用小规模地址枚举验证一遍。4.3 向量化与对齐LDG.STS的节奏共享内存那边理顺以后下一个瓶颈就出现在全局内存到共享内存的数据搬运上。GEMM的搬运特点是块状、连续、量大非常适合向量化加载。我在DeepGEMM里尽量让每一次从全局内存读取都是16字节粒度也就是float4级别的读取这样一次访存能取回4个浮点数访存指令数少了四分之三。这里有个对齐细节向量化读取要求起始地址至少是16字节对齐。如果一段数据恰好从偏移量4字节开始我就没法直接用float4要么手动拆分要么在数据布局上填padding。我在整个项目里对A和B的leading dimension做了对齐约束宁可在尾部多填一些无效数据也不要让主路径出现未对齐访问。向量化和swizzle要配合着做。先从全局内存用float4读入寄存器再按swizzle后的地址把数据写入共享内存这个写入过程也尽量保持向量化。也就是说寄存器里的4个数要能连续写到共享内存的连续4个位置否则就得拆成4次标量写前面省下的访存指令又吐回去一部分。5. 数值精度与边界情况快但必须算得对5.1 低精度输入、高精度累加的分寸DeepGEMM面向深度学习场景输入经常是FP16或BF16但累加器我会坚持用FP32。原因很简单K维经常是4096甚至更大如果一路用FP16累加舍入误差会在长点积里不断累积最终结果可能在十进制第四位就开始飘了。分块累加本身对精度是有帮助的。因为每一小块只做8或16次乘加局部累计误差被限制在小范围内最后再用FP32把各块的累加结果合起来。这个操作等价于做了一个树形归并的简化版比从头到尾一长条点积要稳得多。不过要注意如果块尺寸取得过大比如BK64时一次性累加64项局部误差还是会偏大我一般把BK控制在8到16之间精度和访存开销都能兼顾。5.2 Inf/NaN与异常输入的处理低精度计算里最容易被忽视的是异常值。FP16和BF16的指数范围有限哪怕输入数据本身是合理的FP32数值转成BF16之后也可能溢出成Inf。一旦Inf进入乘累加链FFMA会把Inf一路传播下去最后整个C矩阵的对应行列全部变成Inf或NaN看上去像是算法崩了。我在DeepGEMM里加了一个可选的缩放机制在K维主循环最开始统计当前分块内A和B的最大绝对值如果超过低精度格式能安全表示的范围就整体乘一个缩放因子再进低精度路径。这个操作让训练和推理时偶尔出现的异常值不会瞬间污染整块结果。代价是多了一次数据扫描所以我只在网络量化或梯度出现异常时才开启这个路径。还有一个很容易踩的坑如果输入本身含NaNFFMA的传播行为和标准IEEE语义不完全一样。有些情况下NaN会被当作普通值参与运算结果取决于指令的浮点模式。我的做法是在内核层面加一个开关——如果检测到NaN输入走一条保守的标量路径虽然慢但结果和全FP32版本完全对齐。5.3 非对齐Shape的兜底实际业务里M和N几乎不可能每次都恰好是128的倍数。我对非对齐shape的处理策略是主体部分用高性能Tile内核尾部剩余的若干行或列用一个专门的小Tile内核兜底。小Tile内核不追求极端性能只追求不炸寄存器和正确性。选择阈值时有个经验如果剩余部分占整个矩阵面积的比例小于3%尾核随便写写就行主路径是绝对性能大头如果比例超过15%就要反思是不是主Tile尺寸选得不合适或者干脆把非对齐维度也纳入主路径用mask处理。mask处理的好处是避免两套kernel切换带来的额外同步坏处是predicated FMA会降低主路径效率。我一般以10%作为切换线。这里还想强调一个正确性验证的方法非对齐shape最容易暴露索引bug。我的常规操作是拿一个M130、N66、K17这种边角料尺寸同时跑naive版本和DeepGEMM版本做数值对比误差阈值设到1e-2以内。如果这个形状能过再上大数据集。6. 调优实测从Profiling到参数收敛6.1 先用Performance Metric锁定瓶颈DeepGEMM第一个可用版本跑通后我做的第一件事不是盲调参数而是用硬件性能计数器把三类指标拉出来计算单元利用率、共享内存吞吐、全局内存吞吐。这三者的比例能直接告诉我瓶颈在哪。如果计算单元利用率很高但全局内存吞吐也高说明数据搬运还能和计算重叠得更好如果计算单元利用率不到50%但共享内存吞吐已接近上限那就是bank conflict或swizzle没到位如果指令数看起来很多但不是FFMA占大头那可能是地址计算和数据搬运指令过多该考虑常数下标或减少同步点。拿这些数据说话的好处是不会因为“某次改动感觉变快了”就自我感觉良好。我经常遇到的情况是改了循环展开因子直觉上指令数少了实际因为寄存器溢出导致perf曲线反而下降这种反直觉问题只能靠计数器和profiling才能发现。6.2 一次Tile参数扫描的完整记录有一次我把目标形状定为M512、N4096、K4096的训练场景跑了三组参数结果很有意思参数组BM×BNBKWM×WNTM×TN计算单元利用率实测TFLOPS相对基线A128×128832×328×885%18.2基线B128×1281632×328×888%19.04.4%C256×128864×328×1672%16.0-12%A和B的差异主要在BK上。BK从8翻到16共享内存搬运次数减半访存指令变少利用率提升3个点这个收益是实实在在的。C组看似更激进的Block Tile实际因为每个线程的Thread Tile变成8×16后寄存器压力过大导致occupancy下降计算单元喂不饱数据反而倒退了12%。这次扫描让我明白一个原则参数之间像木桶的木板单独调高一块不一定会变好反而可能挤爆另一块。准确的做法是一组一组调每次只动一个变量记录一版结果再回滚到最佳点做下一个变量。6.3 最终效果与适用边界经过几轮扫描DeepGEMM在目标形状上稳定到了比最初版本快40%左右换句话说从naive到最终版整体提升了一个数量级。但要诚实讲这套参数只在目标形状附近有效。M1的推理场景我会完全换一套参数Block Tile缩到64×64K维推进加大到64主循环里还额外开了mask路径因为这个场景下访存已经是主要矛盾计算强度太低用计算密集型的参数跑就是自找麻烦。这里也顺带提一下DeepGEMM适合什么、不适合什么它适合形状相对固定、需要融合算子、对数值行为有特殊要求的场景它不适合形状极其动态、每次调用都要切换kernel的开销、或者你只是需要一个偶尔跑一次的通用矩阵乘法。项目里我同时保留了naive版本和快速路径就是尊重这个边界。我现在养成的习惯是拿到任何一个GEMM需求先写一版naive作正确性基线再用三层分块的框架套一层快速路径然后永远只在profiling数据指引下动下一步。每调一版参数我就把结果记录到一张表里包括相对性能和当时的硬件指标这样即使两周后回来看也能立刻知道哪个方向已经试过、哪个方向还有潜力。最后再分享一个小技巧调优时留着所有历史版本的kernel切换开关别删旧代码否则你永远无法确认某个性能回退到底是新改动引入的还是硬件状态波动导致的。DeepGEMM这个项目做到后面最大的收获反而不是那百分之几十的性能而是这套“先算账、再分层、靠数据收尾”的优化方法论。