做深度学习加速的朋友,应该都清楚一件事:无论上层网络怎么变,落到计算单元上,最核心、最绕不开的算子之一就是矩阵乘法。DeepGEMM这个名字,听起来是个库,实际上代表了一类专为深度学习场景打造的通用矩阵乘法实现。它的核心使命很直接:在调度和访存上面做文章,把GPU的算力尽量榨干,同时还要兼顾混合精度、融合算子、各种不规则形状这些实际需求。
这篇文章把我自己从零写一个DeepGEMM内核的过程、踩过的坑、以及最后怎么把性能调上去的思路完整讲一遍。适合正在做推理引擎、训练框架底层优化,或者单纯对算子优化感兴趣的开发者。不需要你有很深的经验,但至少要知道线程块、共享内存、矩阵乘法这些基本概念,看起来会更顺手。
1. 为什么要有专门面向深度学习的GEMM
1.1 深度场景里的矩阵乘法到底长什么样
传统的高性能计算矩阵乘法,通常处理的是非常规整的大方阵,一次算完,精度要求高,形状长期一致。但深度学习里的矩阵乘法不一样,它有几个让人头疼的特点。
第一,形状极度不规则。虽然绝大多数场景都是二维矩阵相乘,但M和N往往不是对齐到好尺寸的,比如M等于37、N等于101这种奇怪数字。因为网络里的特征图尺寸、batch大小、序列长度,天然就不考虑底层对齐。
第二,批次维度。深度学习中经常要同时处理多个样本,比如一个batch里有8张图,每张图算一个矩阵乘法。你当然可以把batch维度当成K或者M的一部分去拆,但这么做会引入非常复杂的索引计算。
第三,精度需求是混合的。训练阶段常用FP16、BF16,推理阶段可能用INT8、FP8,甚至INT4。传统GEMM库追求FP64高精度,但深度推理根本不需要这么高的精度,反而更关心吞吐量。
第四,矩阵乘法结果通常不是终点。后面往往跟着偏置、激活、裁剪、残差连接等操作。如果每个操作都单独开一个kernel,结果写回全局内存再读出来,那带宽浪费是灾难性的。
所以,一个为深度学习专门设计的GEMM库,必须针对上面这些特点去优化,而不是简单复用传统科学计算的做法。
1.2 直接调官方BLAS库的烦恼
很多人一说到矩阵乘法,第一反应就是调用官方BLAS库。官方库确实做了大量优化,常规大矩阵性能很好,但它有几个在深度学习场景下的明显短板。
官方库的调度策略是黑盒的。你传一个M、N、K进去,它内部选什么算法、用什么分块、是否切分了batch,你完全不知道。这个黑盒在常规形状下表现不错,一旦遇到小矩阵、非对齐矩阵、或者需要大量重复调用的场景,性能就非常不稳定。
还有一个很现实的问题:形状补齐。官方的GEMM接口往往要求矩阵尺寸对齐到某个倍数,比如8或16。如果M是33,它可能内部给你补到64,多出来的那些行也算了一遍,浪费算力。在小矩阵场景下,这种浪费比例高得惊人。
再就是kernel启动开销。深度学习模型一次推理会执行几十上百个矩阵乘法,如果每个都走一遍完整的库调度逻辑,启动开销会被放大。自己写一个轻量、可预测的调度层,能省掉大量隐形成本。
1.3 DeepGEMM的设计目标拆解
做这样一个库,我一开始就给自己定了几条设计原则。
功能上要支持形状打断。不能假设所有矩阵都是规整的,得用尾块处理机制覆盖所有M、N、K组合。精度上要支持FP16、BF16、INT8、FP8这些常用类型,至少做到能灵活切换。性能上要尽量接近当前主流GPU的理论峰值,不能只满足于“比朴素写法快”。结构上要有清晰的模板化调度,保证编译期知道怎么展开,而不是运行时反复判断。
还有一个很重要但容易被忽略的目标是易集成。库要输出一个简单的API,让上层框架可以直接调用,最好支持零拷贝传入已有的内存布局。如果为了性能让用户做一堆内存重排,那这库再快也很难落地。
2. 核心优化原理,把它拆开看
2.1 朴素写法为什么不行
先看最基础的矩阵乘法写法。
for (int i = 0; i < M; i++) { for (int j = 0; j < N; j++) { float sum = 0; for (int k = 0; k < K; k++) { sum += A[i * K + k] * B[k * N + j]; } C[i * N + j] = sum; } }这段代码从逻辑上没问题,但计算效率非常差。原因是它完全没有利用局部性。内层循环每一次都要读A的一个元素和B的一个元素,这两个元素在内存里距离很远,数据几乎不可能还在缓存里。用通俗点的话说,就像工厂里每个工人都自己跑仓库取一次零件,而不是把一批零件搬到自己工位上慢慢组装。
在GPU上,这种问题会更夸张。因为GPU有几千个线程在跑,如果每个线程都在访问不同的内存地址,缓存命中率会低到让人崩溃。所以矩阵乘法优化的一切,本质上都围绕“如何让计算尽量复用已经取进来的数据”展开。
2.2 分块:第一条主线
解决局部性的标准手段是分块。把大矩阵切成小的子矩阵块,让每个线程块只负责一小块区域的乘法。
假设线程块负责输出一个32×64的C子矩阵。那么它需要A矩阵里对应的32行、B矩阵里对应的64列。如果K循环每次处理一个32×K的A块和一个K×64的B块,这两个块就能先读到共享内存里,后续计算全部从共享内存读,而不是从全局内存读。
共享内存的访问速度比全局内存高一个数量级,而且带宽不受全局内存总线限制。通过分块,把大量本来会重复访问全局内存的操作,压缩成一次加载。这就像批发进货然后囤在本地仓库,再也不用每次零买。
分块参数的选择有讲究。TILE_M和TILE_N越大,数据复用度越高,但共享内存消耗也越大。TILE_K影响K循环的迭代次数和流水深度,太小则无法隐藏内存延迟,太大则共享内存装不下。我一般用128×128的输出块配64的TILE_K,在大多数主流GPU的共享内存容量下都能跑得不错。
2.3 寄存器阻塞和线程映射
分块是线程块级别的策略,但真正让计算跑满的关键是每个线程怎么映射到输出。
一个常见的误区是每个线程只算一个输出点。确实能跑,但算力利用率极低,因为每个线程需要反复发起加载和计算,算术强度太小。正确的做法是寄存器阻塞:每个线程连续计算多个输出点,把数据留在寄存器里复用。
假设一个线程块负责64×64的输出,线程块有8×8共64个线程,那么每个线程可以负责8×8=64个输出元素。这样,B矩阵的一个块被加载后,会被线程重复使用很多次,大幅度减少重复访存。
在GPU上,线程的排列方式也要考虑纵横向。如果线程按行映射,共享内存读取时容易产生bank冲突;按列映射又会影响向量化加载。我经验上是让线程在N方向上连续,这样写入输出矩阵时可以按向量宽度连续写,效果更好。
2.4 流水线:让访问和计算不抢时间
分块解决了复用问题,但K循环仍然会有一个隐藏瓶颈:每次迭代都要先等共享内存加载完,再开始计算。加载和计算串行执行,等于让计算单元一半时间都在等待。
优化方式是软件流水线,也就是双缓冲。给A、B矩阵各准备两份共享内存,一块用于当前K迭代的计算,另一块同时预取下一个K迭代的数据。这样计算单元一直在算,内存系统一直不停加载,两者重叠起来。
双缓冲的实现思路不复杂,但细节很烦。加载和计算之间需要正确的同步,还要防止读写同一块内存造成的数据竞争。我一般把K循环拆成两个阶段:先加载第一份,进入循环后先算旧缓冲区、同时用新缓冲区加载下一块,等两件事都完成再交换角色。
说难听点,这一步骤如果没做好,后续所有调优都会事倍功半。流水线深度拉起来了,性能才能真正起飞。
2.5 张量加速单元的适配
光靠普通乘加指令,即使分块和流水线做到极致,也比不过硬件里的专用矩阵运算单元。这类单元本质上就是一小块专用的矩阵乘法加速器,一次可以执行一个小的矩阵乘加操作,吞吐量远高于普通指令。
但要用好它,需要满足两个条件。
第一,数据布局要对齐。专用矩阵单元的输入通常要求特定的形状和内存布局,比如在内部数据结构里,两个操作数的内存排布必须满足连续性和对齐要求。如果直接传一个普通布局的矩阵,就得先在共享内存里做转置和重排。
第二,精度支持是有限的。专用单元往往不支持FP64,但对FP16、BF16、INT8这类深度学习常用精度做了很大优化。这也是为什么DeepGEMM这类库必然要围绕低精度展开。
我在实现时,对内层最热的计算循环,会做专门的数据打包,确保送入专用矩阵单元的每一批数据都正好卡在对齐边界上。除此之外的循环,比如加载、累加、Epilogue,都是用普通指令完成的。
3. 实际操作:从零写一个可运行的DeepGEMM内核
3.1 确定接口与数据布局
动手写代码前,先定接口。我不会一上来就搞很复杂的API,而是定义一批清晰的配置结构和启动函数,方便后续扩展。
需要考虑的点包括:矩阵A、B、C的步长和偏移,是否带batch维度,是否做了转置,以及累加器的初值。我习惯把这些信息塞进一个描述结构里,避免每个参数都单独传一遍。
布局选择上,主流的深度学习框架通常使用NCHW这种图像布局。对于2D矩阵乘法来说,如果直接把H×W当M,需要先做一次reshape,这一步往往要额外拷贝。为了减少拷贝,DeepGEMM会尽量提供原布局重写的路径,比如对卷积场景直接按矩阵乘法的形状来切分,而不产生中间结果。
举个例子,一个卷积层如果拆成GEMM,插到矩阵里的数据布局往往是完全不连续的。这时候如果硬要用连续内存,就得先做im2col,产生额外的显存消耗。更聪明的方案是在读取阶段就做索引映射,读取时拼出正确的矩阵块。
3.2 朴素参考和正确性验证
优化内核之前,我会先写一个朴素的CPU参考实现,同时用它做正确性的基准。
参考实现非常简单,就是最普通的三重循环。它的作用不是跑得快,而是答案标准。后续GPU内核跑出来的结果,都要和它对一下,允许误差范围和平台有关,但至少要保证大数位一致。
这一步看起来不起眼,实际是优化的安全网。没有这个基准,后面程序一出错,很难判断是调度问题、数据打包问题还是纯粹的算法问题。我见过不少同行在性能调试时被错误结果折磨一整天,最后才发现是最初的数据索引写错了。
验证方法也比较直接,先生成小尺寸数据,比如16×16乘16×16,跑一遍对比。没问题之后,再逐步放大到256×256、512×512,并加入非对齐尺寸和batch维度。每次放大都验证一次,性能调优过程中也保持这个习惯。
3.3 分块内核第一次点亮
接下来是最核心的部分,写一个分块的GPU内核。
下面的代码是一个简化版本,符合GPU编程模型常见的线程块抽象,但省略了不少细节,重点展示思路。
#define TILE_M 64 #define TILE_N 64 #define TILE_K 16 #define THREADS 256 __shared__ float sA[TILE_M * TILE_K]; __shared__ float sB[TILE_K * TILE_N]; __global__ void gemm_v1( const float* A, const float* B, float* C, int M, int N, int K) { int blockRow = blockIdx.y * TILE_M; int blockCol = blockIdx.x * TILE_N; float acc[TILE_M / 8][TILE_N / 8]; // 每个线程负责8x8输出 // 初始化 acc 为 0 for (int k0 = 0; k0 < K; k0 += TILE_K) { // 协作加载 sA 和 sB for (int idx = threadIdx.x; idx < TILE_M * TILE_K; idx += THREADS) { int row = idx / TILE_K; int col = idx % TILE_K; sA[row * TILE_K + col] = A[(blockRow + row) * K + k0 + col]; } for (int idx = threadIdx.x; idx < TILE_K * TILE_N; idx += THREADS) { int row = idx / TILE_N; int col = idx % TILE_N; sB[row * TILE_N + col] = B[(k0 + row) * N + blockCol + col]; } __syncthreads(); // 每个线程计算一个8x8子块 for (int i = 0; i < 8; i++) { for (int j = 0; j < 8; j++) { float sum = acc[i][j]; for (int k = 0; k < TILE_K; k++) { float aVal = sA[(threadIdx.x / 8) * TILE_K + k]; float bVal = sB[k * TILE_N + (threadIdx.x % 8) * 8 + j]; sum += aVal * bVal; } acc[i][j] = sum; } } __syncthreads(); } // 写回 C,这里略过边界判断 for (int i = 0; i < 8; i++) { for (int j = 0; j < 8; j++) { int row = blockRow + (threadIdx.x / 8) + i; int col = blockCol + (threadIdx.x % 8) * 8 + j; C[row * N + col] = acc[i][j]; } } }这个版本跑起来,性能不会很好,但它是第一个能正确工作的分块版本。第一次点亮的意义非常大,它验证了整个分块逻辑、线程调度、共享内存读写路径都是通的。
从这个版本开始,后面的优化是一步步叠加的:先加边界判断和尾块处理,再做双缓冲,再把累加循环展开,最后把数据打包成适合专用矩阵单元的格式。
3.4 K循环预取与双缓冲
分块版本跑通后,我做的第一个大优化就是双缓冲。这一步改动比较大,我直接上一份新的循环骨架。
我在共享内存里为A和B各开两块空间,用parity变量切换当前使用哪块。K循环被拆成两部分:第一次加载先提前完成,然后循环体内先发起下一次预取,再计算当前数据。
__shared__ float sA[2][TILE_M * TILE_K]; __shared__ float sB[2][TILE_K * TILE_N]; int parity = 0; // 预加载第一块 load_tile(0, 0, parity); __syncthreads(); for (int k0 = 0; k0 < K; k0 += TILE_K) { int nextParity = parity ^ 1; if (k0 + TILE_K < K) { load_tile(k0 + TILE_K, 0, nextParity); } compute_tile(k0, parity); __syncthreads(); parity = nextParity; }注意load_tile里面不会立即同步,而是在发起所有加载指令后统一等待。compute_tile使用的数据是已经加载完成的旧缓冲区,其同步已经由上一轮完成。这样计算和加载能够重叠,实际计算单元不会等内存。
这个优化做完后,性能通常会提升一个档次。我见过很多初学者在这个环节卡住,主要原因是同步放错了位置,导致预取的数据覆盖了正在计算的数据。经验是每次都画清楚每个线程不同时刻在访问哪块缓冲区,再来写同步。
3.5 跑通自动化调优
分块参数并不是越大越好,也不是某个值通吃所有形状。我的做法是写一个简单的自动化调优脚本,枚举一组参数组合,每个组合跑一小段基准,记录性能。
常见的参数维度包括:输出块大小、每个线程负责的输出大小、TILE_K、线程块内线程数量、是否使用双缓冲、以及数据打包方式。
| 输出块 | 每线程输出 | TILE_K | 线程数 | 共享内存 | 相对参考性能 |
|---|---|---|---|---|---|
| 64×64 | 4×4 | 16 | 256 | 6KB | 2.1倍 |
| 64×64 | 8×8 | 16 | 64 | 6KB | 2.5倍 |
| 128×64 | 8×8 | 32 | 128 | 10KB | 4.3倍 |
| 128×128 | 8×8 | 32 | 256 | 20KB | 5.1倍 |
| 128×128 | 16×8 | 64 | 128 | 40KB | 5.5倍 |
这张表是我早期在某个GPU实例上的真实记录,主要想说明一个道理:参数是调出来的,不是看论文抄出来的。
自动化调优本身不复杂,我把它设计成两层循环。外层遍历不同配置,内层对一个中等规模的矩阵反复跑几十次,取稳定值。跑完后输出一个性能排行榜,选出最优配置。之后针对特殊形状再做一轮局部搜索,比如固定最优块大小后,微调线程映射方式。
4. 融合、精度和低比特计算
4.1 Epilogue融合为什么价值最大
单独写一个GEMM再加一个激活函数kernel,性能损失主要不在计算,而在于那一次全局内存写和一次全局内存读。一个方阵如果是4096×4096,数据就是64MB的规模,这么来回读一遍的带宽开销非常可观。
融合Epilogue的意思是:矩阵乘法累加的结果不直接写回全局内存,而是留在寄存器或共享内存里,直接完成偏置、缩放、激活等操作再写回。
我最常用的融合形式是:
float alpha = 1.0f, beta = 0.0f; float bias = biasVector[col]; float act = alpha * acc + beta * src + bias; act = act > 0.0f ? act : 0.0f; // ReLU C[row * N + col] = act;这样C矩阵根本不会单独落一次地。看起来简单,节省的带宽非常可观。尤其在推理引擎里,后续马上接下一个算子,减少一次整矩阵访存的影响非常显著。
4.2 BF16、FP16与INT8、FP8的要点
深度学习矩阵乘法对精度的容忍度比较高,所以低精度几乎是必选项。
FP16的优点是范围够用,缺点是表示精度有限,累加容易漂移。BF16的优点是动态范围大,不容易溢出,缺点是尾数精度少。INT8需要额外做量化,要在计算前把浮点数据缩放到整数区间。FP8则更极端,通常配合分块缩放。
我处理这些低精度类型时,有一个不变的原则:累加器始终用FP32。无论操作数是什么类型,乘法结果的累加都用高精度累加器,只在结束阶段再转换回输出精度。这个策略能有效避免几乎所有的数值稳定性问题。
INT8还要小心溢出。两个8位整数乘起来要到16位,如果一次乘加累计很多项,直接用8位重铸就会彻底乱掉。所以INT8的GEMM内核,内部累加寄存器必须是32位整数。
FP8的难点在于动态范围不够。通常需要按块分配缩放因子,每组数据一个scale。这个操作看起来简单,做起来麻烦,因为缩放因子的计算本身要在加载阶段完成,还要在Epilogue阶段恢复,增加了很多指令。
4.3 把GEMM扩展到卷积和注意力场景
很多年前我做卷积加速,第一反应是im2col加GEMM。im2col的好处是把任意卷积转成统一的矩阵乘法,坏处是内存爆炸。一个3×3卷积的输入矩阵可能膨胀9倍。
DeepGEMM里面处理卷积,我更推荐的做法是把卷积拆成几个矩阵乘法的组合,或者使用“隐式GEMM”的思路。所谓隐式,就是不实际构建填充后的矩阵,而是在加载共享内存时直接按感受野索引去取原始数据。这样的话内存占用和效率都更友好。
注意力机制的计算也类似。Q乘K得到注意力分数矩阵,本质就是一个GEMM,后面再接矩阵乘V。但实际计算中还可以做一次算子融合,比如在QK相乘后、softmax之前,把mask和缩放一并处理好。这些融合逻辑完全可以在GEMM的Epilogue阶段完成,不必单独开kernel。
5. 调试、性能分析和一些疑难杂症的排查
5.1 我第一版跑翻车的几个现场
第一次写分块内核时,我遇到最典型的错误是共享内存越界。输出的C子矩阵在边界处超过矩阵范围,访问非法地址,程序直接崩溃。这个问题调试起来特别麻烦,因为崩溃点在访问发生之后很久。
排查方案是先用很小的矩阵跑,矩阵尺寸设为块大小的整数倍,让问题先暴露在逻辑层。确认逻辑没问题后再处理边界。处理边界的方式通常是给每个加载和写入都加上判断,判断非法位置就填充0或者跳过。
另一个坑是bank冲突。共享内存是按bank组织的,如果多个线程同时访问同一个bank的不同地址,就会发生冲突,访问被串行化。早期我为了图省事,把一个矩阵按行连续存放,结果线程按列取数时产生了16路冲突,性能掉了一半还多。
解决bank冲突的标准办法是给共享内存加padding,也就是每行后面多留几个空位。比如把TILE_N从64改成65,这样每一行错开一个bank,冲突就消失了。
5.2 从数据层面定位性能瓶颈
性能不是靠感觉调的,我的做法是先拿性能分析工具看几个关键数字。
第一是寄存器利用率。如果一个线程用了太多寄存器,可能造成占用率下降,影响延迟隐藏。但占用率高也不一定好,可能和内核波浪数相关。第二是共享内存bank冲突率,这个数据能直接指出我加padding的收益。第三是全局内存的命中率,太低说明分块策略失效。
拿到这些数据之后,再画一个roofline模型。横轴是算术强度,纵轴是性能。如果点落在带宽受限区,优化重点是加载和缓存;如果点落在计算受限区,优化重点是减少无效指令和提升展开度。这套方法能避免瞎调。
5.3 常见问题一页速查
我把实际过程中遇到的一批高频问题整理成一个表格,方便大家对照排查。
| 现象 | 可能原因 | 排查与解决办法 |
|---|---|---|
| 内核崩溃 | 共享内存越界或索引计算错误 | 用块大小整数倍的矩阵测试,检查边界判断 |
| 结果数值误差大 | 累加精度丢失 | 累加器改回FP32,避免中途截断 |
| 性能不到预期一半 | 共享内存bank冲突严重 | 给共享内存行加padding,调整线程映射方向 |
| 编译时间极长 | 模板实例过多 | 限制常用形状的模板展开,冷门形状走慢路径 |
| 大量形状性能差 | 尾块处理不当 | 锁定最优块大小后,针对非对齐尺寸单独调优 |
| 集成后与框架结果不一致 | 矩阵布局假设不同 | 先核对A、B、C的步长,确认是否存在转置混淆 |
这个表不是背下来就完事,真正有价值的是每次遇到问题能快速定位方向。我发现自己调试得多了之后,很多问题第一反应不是改代码,而是先跑一遍小矩阵,观察现象模式和表格里哪一行最匹配。
6. 一些额外建议和真实体会
如果让我重新做一遍DeepGEMM,我可能会从更小的问题入手,比如先只支持一种类型、一种形状,把完整流程跑通,再逐步扩展。一开始就想做一个通用高性能库,基本会陷入无止境的参数调优和形状适配,效率不高。
还有一个值得提醒的点是:写GEMM内核不代表一定要从零造轮子。官方的BLAS库在常规大尺寸下很稳定,拿它做性能基准非常合理。但如果遇到库性能不稳或者需要算子融合的场景,自己写一小段专门的内核,性价比相当高。
在实际使用中,我发现融合Epilogue和自动调优是投入产出比最高的两块。强烈建议读者优先把这两件事做好,它们几乎决定了整个库能不能在实际场景中站稳脚跟。
再分享一个小技巧:不要只看单个矩阵乘法的峰值性能,要看完整模型推理的端到端时间。很多时候GEMM本身已经很快了,但前面的reshape、转置、拷贝等数据搬运占了大量时间。把数据通路优化好,比单纯把kernel算得更快更重要。
DeepGEMM这个项目做下来,我个人最大的体会是:高性能计算不神秘,但也不简单。它的门槛不在某个神奇的单点技巧,而在你需要同时理解算法、硬件、内存模型,再把它们拧成一股绳。能坚持把一条主线做透,性能自然就出来了。