
DeepGEMM 这个名字对不搞算子开发的人可能有点陌生。但你只要跑过 Transformer 训练或推理就绕不开它背后的东西——GEMM也就是通用矩阵乘法。市面上大多数高性能矩阵乘算子都被人封装成了黑盒点开文档只有几个参数剩下的全是玄学。我之所以想写 DeepGEMM是因为在一个轻量级推理引擎项目里单纯三层 MLP 的矩阵乘耗时能占到整个网络输出的 60% 以上而换成自己写的一版专用 GEMM 算子后推理时延直接砍掉三成。这个项目算不上什么惊天动地的工程但它把“矩阵乘法为什么慢”“共享内存怎么用”“FP16 的坑在哪里”这些问题一个个榨干了。这篇文章不打算停留在“DeepGEMM 效果很好”这种结论上而是把它的设计思路、核心实现、调参过程和排查经验完整摊开。适合三类人读想搞清楚 GPU 底层并行计算的开发者做推理引擎或算子库优化的工程师以及准备拿矩阵乘法练手性能优化的学生。只要你有一定的并行计算基础哪怕只是写过一点 GPU kernel这篇文章的内容也能直接落到你的项目里。1. 为什么矩阵乘法值得死磕1.1 GEMM 在深度学习里的位置很多人听到“矩阵乘法优化”会下意识觉得这是 HPC 老领域跟深度学习关系不大。但只要你把注意力模型拆开看从注意力分数的计算到前馈网络的两次线性变换全部都是 GEMM。一个 BERT 类模型里矩阵乘的计算量通常能占到总计算量的 85% 以上。换句话说GEMM 快整个网络就快GEMM 慢其他层面再优化也填补不了这个洞。以一层常见的 Transformer 编码器为例设序列长度为 512隐藏维度为 768则 QK^T 的乘法规模是一次 512×768 乘以 768×512前馈网络的两层线性变换也各是一次大矩阵乘法。更直白一点token 越长、模型越大GEMM 的占比就越夸张。所以在业界不管是训练还是推理大家都会盯着矩阵乘算子看峰值算力利用率。一个算子如果能跑到硬件峰值的 70% 以上那整网性能基本就有了保证。1.2 官方库很好但自研理由依然充分看到这里自然有人会问现在主流硬件厂商都提供了官方矩阵库性能已经压榨到了接近峰值为什么还要自己写一个 DeepGEMM我的答案很简单官方库是黑盒而业务经常需要白盒。比如你想在矩阵乘的结果后面立刻跟上激活函数、残差相加或者向量缩放官方库的做法通常是先把 GEMM 结果写回显存再启动下一个 kernel 去处理。这个写回和再读的过程增加了额外的内存带宽开销。而自研算子可以把这些算子融合在同一个 kernel 里数据停留在寄存器或共享内存中就完成了后续操作省掉的访存时间非常可观。此外官方库为了适配各种形状、各种精度、各种硬件内部会有很多分支判断和回退逻辑。当你的业务非常专注比如只用 FP16、固定 block 大小、只跑某种特殊 shape 的时候自研算子可以做更极端的假设去掉大量兼容代码性能上限反而更高。最后还有一个很现实的原因研究 GEMM 优化本质就是研究硬件特性理解了共享内存、寄存器、指令级并行这些底层机制对你做任何深度学习算子优化都只有好处。DeepGEMM 的定位不是要取代官方库而是做一个聚焦于深度学习场景的专用矩阵乘算子默认支持 FP16/BF16 输入、FP32 累加并预留算子融合接口。下面我会把从零到一的过程掰开揉碎讲清楚。2. 设计思路先从硬件账本开始2.1 先算一笔内存账优化 GEMM 如果不算内存账基本等于闭着眼睛开车。这里我先用一个小例子说明问题。假设我们要计算一个 MNK4096 的 FP16 矩阵乘。总计算量是 2×M×N×K也就是大约 137 GFLOPs。如果按最朴素的思路每个线程负责一个输出元素每算一个输出都要从全局内存读 A 的一行和 B 的一列那么 A 和 B 的数据会被反复读取无数次。A 矩阵本身只有 4096×4096×2 字节也就是 32MB但被重复读取 n 次之后实际从全局内存搬进来的流量远远超过 32MB。最终限制性能的根本不是算力而是内存带宽。GEMM 算法层面的优化其实就是提升算术强度也就是“每次从内存取来的数据能参与多少次计算”。理想情况下一块数据从全局内存加载到共享内存会被一个 block 内所有线程复用再从共享内存加载到寄存器又会被同一个线程复用多次。复用的层数越多实际访存量越低性能就越逼近计算峰值。我通常会在优化前先算目标算子的理论算术强度计算量 FLOPs 除以最小必要访存量A、B、C 各读写一次。对于大矩阵这个值可以到几百甚至上千。但朴素实现的算术强度只有个位数这就是为什么矩阵乘优化前后性能能差出一个数量级。2.2 三层分块把数据喂到离计算最近的地方现代 GPU 的存储结构大致可以分成三层全局内存容量大但延迟高共享内存容量小但快得多寄存器最快但数量极其有限。DeepGEMM 的核心思路就是把矩阵乘的计算过程拆成三个层级让数据一层层从“远”到“近”搬运尽量在最近的地方完成计算。Block 级分块一个线程块负责输出矩阵中的一块BM×BN并把对应需要的 A 分块BM×BK和 B 分块BK×BN加载到共享内存。这一步相当于把全局内存的访问次数降到了“每个 block 只读一次”的级别。Warp 级分块线程块内部会分成多个 warp每个 warp 再负责这个 block 输出块中的一个小块。如果计算所需的 A、B 子块可以从共享内存高效读取那么这一步的关键就变成了怎么让多个 warp 不冲突地读共享内存。寄存器级分块每个线程最终负责多个输出元素同时把对应的 A 片段和 B 片段缓存在寄存器里。累加器也始终留在寄存器中只有当整个 block 计算完才对全局内存做一次写入。这样一层层套下来你可能会觉得逻辑很复杂但性能提升非常直观。我从 V0 版本的朴素 kernel 开始一个输出元素对应一个线程性能大约是硬件峰值的 10% 左右做到三层分块后轻松越过 50%。差距就是这么来的。2.3 为什么不用官方库非要自己搭这套框架前面已经提过官方库的融合问题这里再补一个深层原因自研算子的可控性。当你做算子融合、混合精度切换、动态 shape 调度时需要操作的是寄存器、共享内存、指令级排布这些底层细节。官方库的抽象层级很高你很难在它的基础上“塞进”自定义逻辑强行做反而会牺牲性能。DeepGEMM 的早期版本其实是一个非常朴素的矩阵乘 demo后来为了让它在实际推理引擎里可用我又加入了类似形状分桶的调度逻辑对不同大小的矩阵选择不同的分块参数对相同 shape 的调用则直接命中缓存好的 kernel 配置。这些灵活的东西官方库给不了你但它们才是工程上真正拉开差距的地方。为了得到这些自己写一个 DeepGEMM 完全值得。3. 核心实现要点每一个代码决定都有原因3.1 Kernel 主框架K 循环的节奏感DeepGEMM 的 kernel 骨架可以压缩成下面这段看起来很像 CUDA 的 GPU kernel 示例。它不是完整工程代码但把最关键的骨架表达出来了。#define BM 64 #define BN 64 #define BK 32 #define PAD 1 __global__ void deepgemm_kernel(const half* A, const half* B, float* C, int M, int N, int K) { // 共享内存块第二维加 PAD 是为了避免 bank conflict __shared__ half As[BM][BK PAD]; __shared__ half Bs[BK][BN PAD]; int blockRow blockIdx.y * BM; int blockCol blockIdx.x * BN; float accum[TM][TN] {0}; for (int k0 0; k0 K; k0 BK) { // 协作加载把 A 的分块搬进共享内存 load_tile_A(A, As, blockRow, k0, M, K); load_tile_B(B, Bs, k0, blockCol, K, N); __syncthreads(); // 核心计算遍历共享内存中的一小块 K for (int kk 0; kk BK; kk) { for (int i 0; i TM; i) { float a __half2float(As[threadRow][kk]); for (int j 0; j TN; j) { accum[i][j] a * __half2float(Bs[kk][threadCol * TN j]); } } } __syncthreads(); } // 把累加结果写回全局内存 write_C(C, accum, blockRow, blockCol, M, N); }这里有几个关键点需要解释。第一为什么外层循环是 K 维而不是 M 或 N 维因为一个 block 共享内存装不下完整的 A、B 分块必须按 K 方向切步进每次处理一小段 K然后把累加器留在寄存器里不断累积。这样就保证了 C 分块从加载到写回全程不出寄存器。第二__syncthreads()的位置很有讲究。第一个同步是在共享内存写入之后确保所有线程都完成加载才开始计算第二个同步是在计算完之后确保所有线程都读完共享内存下一轮循环才能安全地往里面覆盖新数据。漏掉一个同步轻则算错重则程序崩溃。第三累加器accum一定要是 float即使输入是 FP16。这个后面我会单独展开讲。3.2 共享内存的 bank conflict加一个 padding 就能改变命运共享内存虽然很快但它不是无限带宽的。它内部被划分成多个 bank同一周期内如果多个线程访问不同 bank就能完全并行但如果访问同一个 bank硬件只能把它们串行化这就是 bank conflict。我以前第一次把分块版本写出来后性能一直上不去用性能分析工具一看共享内存相关指令的 bank conflict 达到了夸张的 4 倍惩罚。问题出在共享内存数组的排布上。比如一个 block 里的线程在同一时刻可能访问As[固定行][kk]这一列。因为二维数组按行存储列方向相邻元素的地址间隔是固定的恰好会落在同一个 bank 上于是一整行线程都在抢同一个 bank。解决办法简单得让人惊讶在共享内存数组的第二维加一个元素的 padding。char那种不加这里说的是像As[BM][BK 1]这种。这个 padding 会让原本“对齐”的地址错开相邻线程访问的 bank 岔开bank conflict 就消失了。我带过的同学第一次看到这个改动后性能提升了 20% 多都觉得很神奇其实背后就是硬件存储结构的基本规则。但这里有个容易被忽视的坑加 padding 不能随便破坏向量化加载的对齐要求。如果你用 16 字节的向量加载指令读共享内存第二维是BK 1会导致每个线程的起始地址不再 16 字节对齐反而可能引入新的开销。我的建议是先用标量加载把正确性跑通再考虑向量化两者叠加时单独做验证。3.3 利用专用矩阵计算单元先把数据排对才能让硬件干活现在的 GPU 基本都集成了专门做矩阵乘的计算单元设计初衷就是一条指令完成一个小规模的矩阵乘而不是靠标量乘加指令慢慢累加。想发挥这些计算单元的实力寄存器里数据的放置必须严格遵循硬件定义的规则。不同硬件对 A、B 片段的分布要求不同但核心思路是统一的一组线程通常是一个 warp共同持有一块 A 子矩阵和一块 B 子矩阵通过特定指令完成批量乘加。DeepGEMM 在早期版本中只用了普通的乘加指令性能到 60% 附近就封顶了。后来我做了一个优化版本让每个线程持有 4×4 的累加器块同时把 A、B 片段按矩阵计算单元的输入要求重新排列再切换到矩阵指令峰值利用率一下子提升了 20 个百分点。如果你在写这类算子我的建议是先把普通乘加指令版本跑通性能数据记录下来再去做矩阵指令版本。这样你能清楚地知道收益到底来自矩阵指令本身还是来自数据布局的改进。不要一上来就照着官方模板库抄矩阵指令的装载方式那些代码为了通用性做了大量抽象初学者很容易迷失。你先试着用最朴素的方式把一个 warp 的输入对齐到一张映射表理解哪个线程负责哪个输出元素之后再套用硬件指令就顺理成章了。3.4 数值精度FP16 输入FP32 累加用深度学习的半精度数据训练或推理最大的印象就是“快”但很少人注意精度陷阱。FP16 的表示范围有限直接拿 FP16 累加很容易在小数值累加时发生溢出或精度损失。BF16 虽然指数范围好一些但尾数位更少逐项相加的误差会更大。所以 DeepGEMM 的做法是输入保留 FP16 或 BF16 数据加载到寄存器后立即转成 float所有的乘法和累加都在 float 下完成。这样虽然多了一次类型转换指令但在现代 GPU 上成本很低而精度收益非常明显。如果项目中需要对比误差你可以分别跑一版 FP32 累加和一版 FP16 累加用同一份随机输入算最大绝对误差结果通常差好几个数量级。这一点在实现 attention 类算子时尤为重要因为 softmax 分母累加和注意力加权求和的中间结果范围差异很大一旦累加精度不足最终输出就可能出现明显的质量劣化。DeepGEMM 从一开始就坚持 FP32 累加不是保守而是实测下来的必经之路。4. 实操过程与性能调优记录4.1 从朴素版本到 DeepGEMM 的四个阶段这部分我想用表格把性能演进的路径拉出来方便你对每种优化手段的收益有直观认知。下面的数据基于我用的那张测试卡理论 FP16 峰值取一个常见的整数实际数字不重要重点是看优化阶梯和波动原因。版本主要优化实测性能TFLOPS硬件峰值利用率关键瓶颈V0朴素版本一线程一元素约 8约 10%全局内存带宽饱和无复用V1共享内存分块BM/BN/BK约 25约 25%共享内存 bank conflict 严重V2寄存器多累加器 向量化加载约 45约 45%指令发射效率不足乘法指令过多V3padding 消除 bank conflict 预取约 55约 55%普通乘加指令达到瓶颈V4切换专用矩阵计算指令约 75约 75%接近硬件上限还有提升空间V0 到 V1 为什么提升最大因为共享内存分块把全局内存访问次数从“每个元素读两次”降到了“每个 block 只读一次”直接解决了内存带宽瓶颈。V1 到 V2 的提升更多来自寄存器复用一个线程负责多个输出元素加载进寄存器的一个 A 值可以被多个输出计算复用减少了共享内存读次数。V2 到 V3 比较“闷”提升不明显但稳定性好了很多bank conflict 的惩罚在复杂 shape 下会迅速放大V3 的收益在 shape 变化时才会体现。V4 就是结硬寨打呆仗用硬件专用指令换普通乘加指令计算吞吐直接翻了一截。4.2 分块参数不是拍脑袋拍出来的BM、BN、BK 这些参数看着像魔法数字其实每一组都有资源约束。GPU 每个线程块能用的共享内存有限比如假设是 96KB每个线程能用的寄存器数量也有限比如 255 个。BK 越大意味着每个 K 切片能缓存更多数据循环次数更少但共享内存占用也线性增长留给其他资源的空间就少。BM 和 BN 同理它们决定了每个 block 要多少共享内存也决定了线程数。我给出一个非常粗的预算方法。共享内存要放两个矩阵分块容量大约是(BM*BK BK*BN) * 2字节FP16或* 4字节FP32。这个值必须小于硬件的共享内存上限同时还要留一点余量给其他用途。寄存器方面每个线程有 TM×TN 个浮点累加器再加上加载 A、B 片段的临时寄存器如果 TM×TN 超过 16 或 20线程数就只能降下来占用率低会导致延迟无法掩盖。选参数时我一般先定 BMBN64BK16 起步跑通后逐步把 BK 加到 32、64。不要一上来就把 BK 拉满寄存器压力和共享内存压力同时上升性能会不升反降。你在自己的机器上做实验时把每个候选参数组都跑一遍记录性能和带宽利用率最后你会发现最优参数往往是让线程块数量和每个 block 的资源占用量刚好平衡的那组。4.3 性能度量的正确姿势衡量 GEMM 算子性能最常用的指标是有效 TFLOPs计算公式是2*M*N*K / 运行时间。但运行时间怎么取很有讲究。GPU kernel 是异步执行的你不能在 kernel 启动后用 CPU 的时钟直接掐表必须在设备端插入事件记录或者用专门的分析工具。我见过不少新手直接用 Python 接口计时结果把 kernel 排队时间和传输时间都算进去了性能数字惨不忍睹。另外不要只用一个 shape 的成绩来夸一个算子。DeepGEMM 在 MNK4096 这种“舒适区”很好看但换到 M1 的推理场景再好的分块也救不了矩阵太小带来的启动开销。所以我在项目里为不同 shape 准备了不同 kernel 配置还做了一个简单的自动选择逻辑。当 M 或 N 小于某个阈值时就切换到“合适即停”的小 block 版本避免一个 4096 规模的 kernel 去处理 128 维的小矩阵。跑性能对比时我习惯每个配置跑至少 10 次去掉最高最低后取中位数然后和硬件理论峰值比一比。如果某个版本的利用率长期低于 50%优先怀疑访存或者 bank conflict如果能达到 70% 以上说明数据路径基本健康剩下的提升只能靠指令级优化了。5. 常见问题与排查心得5.1 问题速查表做算子优化最容易遇到的不是算法难而是排查难。这里整理一个速查表都是我在 DeepGEMM 开发和测试中实际见过的坑。现象可能原因排查方法计算结果全是随机乱数共享内存同步缺失或者数组越界覆盖检查每次加载后的__syncthreads()用计算边界条件做小矩阵验证结果出现 NaN 或 InfFP16 累加溢出或输入数据本身异常确认累加器用的是 FP32再用小数值输入复测性能始终上不去bank conflict、占用率过低、没有向量化用性能分析工具查看 shared memory bank conflict 次数和 warp 占用率kernel 直接崩溃索引越界block 维度配置错误在 kernel 入口加边界判断先用 1×1 线程块跑最小用例换一个 shape 性能差距巨大分块参数不适应当前 shape做 shape 分桶把不同矩阵规模映射到不同 kernel 配置5.2 一个被我低估的细节向量化加载很多人写矩阵乘时共享内存加载是从一个线程加载一个元素开始的。这样写最简单但效率很低。GPU 上一条向量加载指令最多可以一次搬运 128 位数据也就是 8 个 FP16 或 4 个 FP32。换句话说如果一个线程一次只搬一个 FP16相当于浪费了 7/8 的加载带宽。DeepGEMM 在 V2 阶段把全局内存到共享内存的加载改成了向量化方式边界处用标量处理。修改后内存吞吐和指令数都有了明显改善。但这里有个很隐蔽的坑向量化加载要求地址 16 字节对齐。A 和 B 的全局内存指针来自上层框架通常已经是 256 字节对齐的问题不大但共享内存里的地址对齐会被我之前说的 padding 策略破坏。我在第一次尝试“padding 向量化”的时候掉进过这个坑为了消 bank conflict 加了 padding结果向量化加载的对齐乱了性能反而倒退。后来我把 padding 移到数组末尾或者让 padding 也保持 16 字节对齐才同时拿到两者的收益。这类交叉影响的细节不做实测很难意识到。5.3 先正确再优化每次只改一个变量这是我做 DeepGEMM 最大的经验。早期我贪心总想把分块、向量化、矩阵指令一次性全部加上结果某个东西改错了性能和分析工具指向的根因完全对不上排查花了三倍时间。后来我给自己定了个规矩一个版本只改一个变量。比如 V1 到 V2我只把“一线程一输出”改成“一线程多输出”其他全部保持不变。跑出来的性能对比才真正归因到寄存器复用。再比如 V2 到 V3只有 padding 变了。即使性能变化不大我也能明确知道这个 padding 对当前配置的影响。这个习惯在调分块参数时更重要。参数之间会互相影响比如 BK 增大会降低共享内存压力但增加寄存器压力你只有每次只动一个维度才能看清因果。正确性验证也要跑在前头。不要一上来就用 4096 的大矩阵测先用 8×8、16×16 这种小矩阵和 CPU 上的简单实现逐元素对比通过了再慢慢放大 shape。矩阵乘的 bug 往往出现在边界和分块取整处小矩阵更容易暴露索引错误。等小矩阵全对了再用大矩阵冲性能你会省掉大量无意义的 debug 时间。5.4 不要盲目模仿官方库的极致版最后分享一个心态上的建议。我看到过很多同学打开官方开源的矩阵乘模板库一看到那些复杂的调度和内存排布立刻觉得自己写不出来然后放弃了自己动手的计划。其实那些极致版本是为各种硬件和各类形状做通用优化的里面很多复杂逻辑在固定场景下根本用不到。DeepGEMM 的代码比它们简单得多但在特定 shape 下依然能跑到接近硬件上限。我的看法是第一次做算子优化应该以“自己能完整解释每一个选择”为目标而不是以“逼近官方库的每一条指令”为目标。等你把分块、同步、bank conflict、向量化这些基础问题都亲手趟一遍后再回头看复杂模板库的代码你会发现它们的设计意图其实你都能看懂了。到那时候你的 DeepGEMM 才真正成为你自己的东西。如果你准备动手写一个类似的项目最后送一条我从 DeepGEMM 里总结出来的实操建议先把小矩阵跑通再用性能分析工具找瓶颈每次只动一个变量让数据告诉你下一步该做什么。坚持三个版本迭代之后你会回来感谢这份耐心的。