做深度学习部署这些年我几乎每天都要和GEMM通用矩阵乘法打交道。刚开始写算子时我以为把三層循环写对就算完事直到用Profiler一看才发现手写版连硬件峰值算力的5%都跑不到。后来我仔细研究了一个叫DeepGEMM的高性能矩阵乘法项目才彻底明白GEMM优化到底在优化什么。这篇文章就把我从能跑到跑得快整个过程中的理解、代码和踩坑记录整理出来覆盖分块、寄存器缓存、TensorCore、双缓冲、数据布局和量化这几个层级希望对正在做算子优化或者好奇深度学习底层加速原理的朋友有帮助。1. 为什么所有深度学习框架都把GEMM当命根子算力开销的真实构成1.1 卷积、全连接和Transformer注意力都能落到GEMM很多人对GEMM的印象停留在两个矩阵相乘。实际上深度学习模型里的卷积、全连接、注意力机制甚至部分归一化操作最终都能转化成一个或几个大矩阵乘法。拿卷积举例。常规卷积操作是滑动窗口逐个点乘求和看起来很复杂但通过im2col——也就是把输入特征图按卷积核感受野重新排列成一个大矩阵——卷积就变成了一个GEMM。虽然im2col会带来额外的内存占用但换来的却是极高的计算密度和成熟的矩阵乘优化库支持。全连接层就更直接了本质就是矩阵乘加偏置。Transformer里最耗时的注意力计算Q乘K转置、注意力分数乘V更是原封不动的GEMM。我在实际优化某个图像处理模型时统计过各算子的耗时占比GEMM类算子包括被展开的卷积占了总计算量的60%到80%。这个比例在语言模型里更夸张几乎所有的计算都发生在矩阵乘法里。所以GEMM性能每提升10%整个模型的训练和推理时间就能明显下降。这也是为什么各大硬件平台都专门为GEMM准备了专用的矩阵指令单元。1.2 性能天花板由GEMM决定一个直观的数字对比我随便举一个直观例子。假设要计算两个4096乘4096的FP32矩阵相乘总的浮点运算次数大约为4096 × 4096 × 4096 × 2 ≈ 137.4 GFLOPs如果某张GPU的FP32峰值算力是每秒30TFLOPS理论上完成这个计算只需要大约4.6毫秒。但如果用最朴素的单线程循环跑假设每秒只能执行50亿次浮点运算那就要27秒以上。差距接近四个数量级。有人可能会说那我用多线程并行不就行了GPU本来就是大规模并行架构这确实是突破口。但只靠线程多还不够因为矩阵乘法的主要瓶颈往往不在算得不够快而在数据喂得太慢。我实测过一份代码只改线程并行策略算力利用率也就从2%提到10%左右真正让它飙升到70%以上的是后续要讲的分块和寄存器缓存。这就是DeepGEMM这类项目存在的意义它不是简单地实现矩阵相乘这个数学功能而是要把硬件上每一份算力、每一字节带宽都榨干。理解了这一点再去看它的代码和设计就会有完全不同的视角。2. 朴素矩阵乘法为什么慢到让人抓狂DeepGEMM要捅破的三层窗户纸2.1 计算与访存的鸿沟峰值算力远大于内存带宽先建立一个基础认知CPU或GPU芯片内部的乘法器运算速度通常比它读写显存或内存的速度快几十倍。换成生活化的说法算力就像工厂的生产线速度内存带宽就像原料运输车队的速度。如果生产线每秒钟能加工100个零件但运输队每秒只能送来10个生产线就得空转90%的时间。对于朴素矩阵乘法每算一个输出元素C[i][j]就需要读取A的第i行和B的第j列共2N个数据做完N次乘加。于是总访存量大约是2 × N × N × N 2N³而总计算量恰好也是2N³ FLOPs。也就是说朴素实现里每做一次浮点运算就要搬运两个数据。现代硬件的计算访存比单位字节能支撑的浮点运算次数动辄几十甚至上百朴素实现显然完全匹配不上。DeepGEMM优化的第一步就是提高数据的复用率让一个数据从内存读出来后能够参与尽可能多次的计算。2.2 线程并行与数据复用的矛盾加速倍数不够看GPU有几千个计算核心听起来很吓人但核心之间并不完全独立。它们通过线程块和共享内存来协作共享内存的容量和带宽都是有限的。如果每个线程都独立去读全局内存不仅带宽被占满而且大量数据会被重复读取造成严重的浪费。正确的思路是让一个线程组共同负责一个输出子块。比如一个线程块计算64乘64的输出区域需要读取A的64行乘以某个宽度的子块以及B的某个宽度乘以64列的子块。这些数据先被加载到共享内存或寄存器里再被组内所有线程反复使用。这样一来每个数据被读取一次就可以参与64次甚至更多的计算算力利用率自然就上去了。不过这里有个容易忽略的点线程块内部线程之间的合作需要同步同步太多会拖慢速度同步太少又会导致数据依赖出错。怎么平衡就是调度设计的艺术。DeepGEMM里的分块尺寸和循环顺序全都是围绕这个平衡来定的。2.3 精度与格式的博弈从FP32到FP16/TF32/FP8除了数据和计算精度格式也是影响GEMM性能的关键。FP32需要32位存储计算单元也更大FP16和BF16只要16位TensorCore处理起来速度是FP32的好几倍FP8更是把存储和计算开销进一步砍半在最新硬件上能获得数倍于FP16的吞吐。但精度降低会带来数值误差。DeepGEMM这类高质量实现通常提供多种精度策略训练前中期用FP32或TF32做累加保证稳定性推理阶段适当使用FP16或FP8来换取吞吐量。我自己的经验是FP8量化必须配合逐层或逐通道的缩放因子否则模型精度很容易崩。这也是GEMM库看起来只是乘法但实际工程量大得吓人的原因。3. 手写一个DeepGEMM风格的核心Tiling策略、寄存器缓存和Bank Conflict避让3.1 基本分块把大矩阵切成能被缓存装下的小块现在进入实战。我们先不直接上TensorCore而是先看一个基于普通计算核心的高效GEMM怎么写。这里最核心的词汇是Tiling也就是分块。为什么要分块因为缓存和共享内存都很小。如果整个矩阵一次性加载64KB的共享内存根本装不下。如果把输出矩阵切成若干个边长64或128的小块每个线程块只负责一小块那么它所需的A和B子矩阵就能完整放进共享内存。分块尺寸的选择我推荐从64×64开始试。块太大共享内存装不下一个SM上能同时运行的线程块数量变少占用率下降块太小数据复用率不够访存压力大。我在某块GPU上试过32、64、128三种Tile64×64的综合效果最好算力利用率比32提升约15%比128也略高一点原因是128的Tile让寄存器压力明显变大占用率掉了一些。3.2 寄存器缓存与片上内存的配合这里还有个很多人容易忽略的细节共享内存只是中转站真正最快的存储是寄存器。理想状态下A和B的数据从全局内存读到共享内存再从共享内存读到寄存器寄存器里的数据直接参与乘加运算这样每一次访存都势能最大化。为了做到这一点循环的内层通常这样安排每个线程负责输出子块中的一小片元素例如4×4。线程先从共享内存把对应的A片段和B片段取到寄存器数组里然后在最内层循环连续做16次乘加期间不再访问共享内存。这个技巧通常叫寄存器缓存Register Caching。相比每算一个输出元素都去读共享内存寄存器缓存版本能减少约4到8倍的共享内存访问量。配合这个思路我用伪代码整理了一个核心循环骨架// 每个线程负责计算 output[4][4] 大小的子块 float acc[4][4] {0}; // 沿 K 维循环每次处理 kStep 个元素 for (int k0 0; k0 K; k0 kStep) { // 将 A 和 B 的子块从全局内存拷贝到共享内存 loadToSharedMemory(A, B, tileA, tileB, k0); __syncthreads(); // 将共享内存数据分散加载到寄存器 float aReg[4], bReg[4]; for (int subK 0; subK kStep; subK) { for (int i 0; i 4; i) { aReg[i] tileA[threadRow i][subK]; } for (int j 0; j 4; j) { bReg[j] tileB[subK][threadCol j]; } // 内层乘加完全使用寄存器 for (int i 0; i 4; i) { for (int j 0; j 4; j) { acc[i][j] aReg[i] * bReg[j]; } } } __syncthreads(); }这段代码看起来简单但里面的两个__syncthreads()是必须的第一个确保子块完整加载完再被读取第二个确保所有线程都读完了当前子块才能覆盖共享内存准备下一轮。3.3 共享内存Bank Conflict的产生逻辑与规避方法共享内存虽然比全局内存快很多但它也有一个隐藏陷阱叫Bank Conflict。共享内存在硬件上被分成32个Bank每个Bank同一时刻只能服务一个读请求。如果同一周期内多个线程访问的地址落在了同一个Bank的不同地址上就会产生冲突硬件只能串行处理性能瞬间打折。举个最简单的例子如果所有线程都去读共享内存的第0个元素这反而没问题因为硬件支持广播但如果是线程0读地址0线程1读地址32线程2读地址64虽然地址不同但它们都落在同一个Bank上就会触发冲突。规避方法最常见的叫padding。比如共享内存数组原本是float smem[32][32]这时第0行和第1行相同列的元素正好落在同一个Bank模式里。我把声明改成float smem[32][32 1]每行多垫一个float行与行之间就错开了Bank Conflict大幅减少。我测过这个改动效果非常明显同样计算量下仅这一个改动就能带来约20%到30%的性能提升。这也是DeepGEMM类实现里看着很无厘头却又极其关键的小细节。3.4 内存对齐和大页访存容易忽略的另一个性能开关除了共享内存全局内存的访问模式同样值得关注。GPU访问全局内存时如果同一批线程读取的地址能落在连续的128字节或更长的对齐区间内硬件就能用一次事务完成传输反之事务数会增加好几倍。所以加载A和B的子块时最好保证每个线程读取连续的一段数据而不是每个线程跨行读取。常见的做法是让线程ID直接映射到连续的内存偏移上并用向量化加载指令例如一次性读float416字节。我在实现中把加载代码改成float4格式后带宽利用率从40%提升到了80%以上。4. 深度优化三板斧TensorCore、按K流水线预取、量化感知的数据布局4.1 从标量乘加到矩阵指令TensorCore为什么能一个指令算一个子矩阵如果只停留在标量乘加层面即使Tiling做得再好也很难接近新硬件的峰值。因为现代GPU普遍配备了专用的矩阵运算单元通常叫TensorCore。这种单元的特点是一条指令就能完成一个4×4、8×8甚至16×16的矩阵乘累加远非普通核心一条指令一次乘加可比。我第一次看到TensorCore的吞吐数据时很震惊在FP16精度下TensorCore的峰值算力大约是普通FP32核心的4到8倍。也就是说同样跑一个大GEMM用TensorCore天然就有好几倍的性能优势。DeepGEMM当然不会放过这个关键特性。它的核心循环大量使用矩阵乘累加指令让硬件直接在一个时钟周期里完成整块子矩阵的计算。4.2 按K切分与双缓冲把访存和计算重叠起来沿着K维度也就是两个矩阵相乘的内维把大矩阵切成多段是另一个重要的优化。设想一个输出的64×64子块需要累积K维上的所有数据。如果一次性把整个K维的A和B片段都加载进共享内存容量肯定不够。所以标准做法是沿K切成若干小步长每轮只加载一小段。但这里有个隐患每一轮加载完都要等计算完成才能加载下一轮GPU那点空闲时间全浪费了。解决方法是双缓冲在计算第t个K段的同时预加载第t1个K段到另一块共享内存区域。这样访存和计算被完美重叠起来只要预加载时间不超过计算时间理论上可以把全局内存延迟完全隐藏掉。实际操作中我建议把共享内存声明成两份例如__shared__ float smemA[2][blockSize][kStep]; __shared__ float smemB[2][kStep][blockSize];每一轮轮换使用buffer 0和buffer 1。配合异步拷贝指令可以让DMA引擎在后台搬运数据计算单元只管算数。这个改动在我的实测中让GEMM核函数从算力利用率的45%直接提升到了70%以上是所有优化手段里收益最大的一项。4.3 FP8数据布局为什么不是单纯把8位整数对齐就行FP8是最近两年很热门的精度格式它用8位表示浮点数分为E4M3和E5M2两种变体。E4M3精度高一些适合权重和激活E5M2动态范围大一些适合误差累积和梯度场景。很多人以为FP8优化就是简单地把FP16换成FP8实际上远没那么简单。首先FP8需要考虑缩放因子。因为8位的表示范围很有限矩阵中的数值范围如果太大直接截断会损失大量信息。一般用分块缩放即每个32×32或64×64的子块乘一个缩放系数把数值映射到合适范围。缩放系数的选取本身又需要额外的计算和内存访问。其次数据布局要为TensorCore的加载方式量身定制。常见做法是把K维数据重排成特定的分块格式例如按16×16的小块连续存储这样一条加载指令就能把参与矩阵指令的整块数据搬进寄存器。我一开始没在意布局直接按行主序存FP8数据结果TensorCore利用率只有30%。后来换成分组重排布局利用率才恢复正常。我不建议在所有场景里无脑上FP8。如果模型本身对精度很敏感或者没有做校准就会遇到精度崩塌。保守做法是先用FP16把整个GEMM链路调通再在验证集上逐步切换FP8最后再对缩放因子做校准。5. 从跑通到跑快Profiler读法、三处真实翻车记录和一组性能对照5.1 性能分析第一课看算力利用率而不是只看耗时优化GEMM时如果只盯着总耗时容易陷入盲调。我建议先打开Profiler看两个指标实际算力比如TFLOPS和算力利用率实测算力除以峰值算力。只有算力利用率才真实反映GEMM核函数写得好不好。打个比方同一个GEMM在一张峰值50TFLOPS的显卡上跑到20TFLOPS利用率是40%换到一张峰值100TFLOPS的显卡上跑到35TFLOPS看起来绝对速度变快了但利用率只有35%说明核函数在新卡上并没有发挥出应有水平。我的经验是普通实现利用率在10%以下很正常经过Tiling和寄存器缓存能到40%到50%加了双缓冲和TensorCore之后可以冲击70%到85%再往上就需要非常精细的调度和硬件特性配合。另一个需要盯的指标是访存吞吐。如果计算指令一直处于等待状态可能不是算力不够而是数据加载太慢。这时候要去查全局内存的load效率、共享内存的bank冲突次数以及缓存命中率。5.2 三处真实翻车记录Swizzle、Occupancy和精度误差翻车记录一Swizzle实现错了。Swizzle是一种把共享内存地址的XOR重排目的是让同一行内不同线程访问的Bank均匀分散。我一开始只做了简单的重排列结果32个线程访问地址时反而全部落在同一个Bank上性能比我直接按行存储还差。后来我把硬件Bank的数量和每个Bank的宽度画成图逐一对齐才调好。自己写Swizzle之前务必先在纸上推算BankIndex否则很容易适得其反。翻车记录二Occupancy被寄存器数量卡住。我为了让内层循环减少访存把每个线程要用的A和B片段都展开成很多寄存器变量结果每个线程用了超过128个寄存器导致一个SM上能运行的线程块数量骤降。虽然单线程计算快了但并行度严重下降总体性能反而倒退了。后来我通过限制Tile为4×4把寄存器数控制到80以内性能才回来。寄存器不是越多越好它是用并行度换来的。翻车记录三FP8误差在中途才爆发。我在一个视觉模型上切了FP8训练前几步完全正常loss曲线也漂亮到了第几千步直接发散。排查了很久才发现是一个子块的缩放因子只用了16位保存溢出后变成无穷大。将缩放因子提升为32位累加之后稳定性才恢复。这里也让我养成了习惯任何量化后的GEMM必须跑足够长的步数来验证不能只看前几百步。5.3 一组可参考的优化效果对照表我把自己在某块中端GPU上、对一个4096×4096的FP16 GEMM做性能优化的数据整理成表格方便大家对比不同阶段的收益优化阶段主要改动算力利用率相对朴素实现加速比朴素实现三重循环无分块约3%1.0第一轮Tiling64×64分块共享内存约18%5.2第二轮Tiling寄存器缓存Bank Conflict规避约35%9.8TensorCore适配改为矩阵乘累加指令约62%17.5双缓冲异步加载访存与计算重叠约76%21.3数据布局和Swizzle按K重排Bank索引优化约81%22.7注意这个表格只是参考实际数字会因硬件、编译器版本和矩阵形状而波动。但它至少说明了两个趋势第一分块和寄存器缓存能为后续优化打好地基跳过这一步直接上TensorCore收益会大打折扣第二当性能提升到一定程度后每再上一个点都需要处理细节的精细度呈指数级上升。我在实际项目中最后用到的方案往往不是某个单一优化而是上面所有技术的组合。DeepGEMM这类项目的价值恰恰在于它把这条完整的优化路径完整地展现了出来。如果你手头正好有要优化的GEMM场景我的建议是先跑通基线把分块和寄存器缓存做好再加上双缓冲最后才考虑精度压缩和高级Swizzle一步一个脚印性能和稳定性都能兼顾。