☰
DeepGEMM 实战:从零到一编写高性能矩阵乘法算子
2026/10/10 9:27:01 网站建设 项目流程

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 如果不算内存账,基本等于闭着眼睛开车。这里我先用一个小例子说明问题。

假设我们要计算一个 M=N=K=4096 的 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,线程数就只能降下来,占用率低会导致延迟无法掩盖。

选参数时我一般先定 BM=BN=64,BK=16 起步,跑通后逐步把 BK 加到 32、64。不要一上来就把 BK 拉满,寄存器压力和共享内存压力同时上升,性能会不升反降。你在自己的机器上做实验时,把每个候选参数组都跑一遍,记录性能和带宽利用率,最后你会发现,最优参数往往是让线程块数量和每个 block 的资源占用量刚好平衡的那组。

4.3 性能度量的正确姿势

衡量 GEMM 算子性能最常用的指标是有效 TFLOPs,计算公式是2*M*N*K / 运行时间。但运行时间怎么取很有讲究。GPU kernel 是异步执行的,你不能在 kernel 启动后用 CPU 的时钟直接掐表,必须在设备端插入事件记录,或者用专门的分析工具。我见过不少新手直接用 Python 接口计时,结果把 kernel 排队时间和传输时间都算进去了,性能数字惨不忍睹。

另外,不要只用一个 shape 的成绩来夸一个算子。DeepGEMM 在 M=N=K=4096 这种“舒适区”很好看,但换到 M=1 的推理场景,再好的分块也救不了矩阵太小带来的启动开销。所以我在项目里为不同 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 里总结出来的实操建议:先把小矩阵跑通,再用性能分析工具找瓶颈,每次只动一个变量,让数据告诉你下一步该做什么。坚持三个版本迭代之后,你会回来感谢这份耐心的。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询