☰
DeepGEMM:深度学习高性能矩阵乘法内核优化实践
2026/10/11 9:12:29 网站建设 项目流程

GEMM(通用矩阵乘法)是深度学习训练和推理里最底层的重型计算单元,从全连接层到注意力分数计算,本质都是在搬矩阵。很多人以为直接调官方闭源数学库就万事大吉,但真到了性能优化、算子融合、低精度推理这些阶段,手边能自己掌控的矩阵乘法内核反而不够用。DeepGEMM 就是我做的一个面向深度学习场景的高性能矩阵乘法内核集合,目的很直接:在常见形状上把算力吃透,同时把 epilogue 融合、量化支持这些推理引擎真正需要的功能做进去。这篇内容适合正在做推理引擎、写自定义算子、或者想搞懂矩阵乘法性能瓶颈的朋友,我尽量把从分块策略到 Tensor Core 指令、再到精度对齐的整个思路讲清楚。

1. 为什么深度学习场景需要一份"专属GEMM",而不是直接调官方库

1.1 GEMM 是深度学习的发动机,再快也不嫌快

先说一句很多人可能不太在意的背景知识。深度学习模型跑一次推理,计算量绝大部分落在卷积和矩阵乘法上,而卷积在底层实现时也往往会转换成矩阵乘法来处理。所以 GEMM 的性能直接决定了一次前向推理的快慢。

GEMM 的数学公式非常简单:C = alpha * A * B + beta * C。A 是 M×K 矩阵,B 是 K×N 矩阵,C 是 M×N 输出矩阵。看起来就是三重循环的事,但问题在于数据量远超芯片的片上存储。一颗现代 GPU 有几百 GB/s 甚至几 TB/s 的显存带宽,但一个 4096×4096 的 FP16 矩阵就有 32MB,整个模型矩阵动辄几十个这样的块,不可能全部塞进片上缓存。所以 GEMM 优化的核心从来不是"怎么算",而是"怎么搬数据搬得少、算得密"。

我选择自己写 DeepGEMM,不是因为官方闭源库不行,而是因为它解决不了我的三个问题。第一,融合算子。推理引擎里 GEMM 后面基本都跟着 bias、激活函数、LayerNorm、量化缩放这些操作,官方库只负责输出 C 矩阵,剩下的操作要我重新起一个 kernel 读写一遍数据,显存带宽全浪费了。第二,形状控制。官方库对超大连续矩阵调校得很好,但碰到 M=1 的推理场景、或者 N 比较小的窄矩阵,性能优势就没那么明显了。第三,可控性。我想针对自己的模型形状做定制,想在看到性能瓶颈时能精确知道是哪一行代码在等待,闭源库做不到。

1.2 DeepGEMM 的定位:不是重新造轮子,是造一套能改的轮子

DeepGEMM 的定位是一套面向深度学习推理场景的 GEMM 内核模板集,不是要替代官方库在所有场合的工作,而是重点覆盖推理引擎里最常用到的几种情况:常见批量大小下的大矩阵乘、低精度输入(FP16/BF16,甚至量化后的 INT8/FP8)、以及需要把额外算子焊进 epilogue 的融合需求。

做这套东西之前,我先明确了一个原则:先让内核在一个小形状上跑出明显的性能,再考虑泛化,而不是一开始就试图处理所有边界情况。我最初只实现了 M=4096、N=4096、K=4096 的基础版本,跑通之后再逐步往外扩。这样做的原因是 GEMM 优化的变量实在太多,tile 尺寸、寄存器布局、流水线深度、访存模式,任何一个改动都可能改变性能特征,如果不固定形状去调参,很难知道到底是哪个改动起了作用。

2. 分块调度:把大矩阵切成能塞进芯片的小方块

2.1 从数学公式到三层存储层级

矩阵乘法最直观的实现就是三重循环累加,但这样每个 A 元素会被读 N 次,每个 B 元素会被读 M 次,数据搬运量惊人。分块的目标是提高数据复用:A 矩阵的某一行会被 C 矩阵同一行的所有列使用,B 矩阵的某一列会被 C 矩阵同一列的所有行使用,所以把计算切成 M_BLOCK × N_BLOCK 的小方块后,这个小方块计算只需要加载对应的 M_BLOCK×K 的 A 分块和 K×N_BLOCK 的 B 分块,数据复用率从 1 提升到了块尺寸级别。

在 GPU 上,分块要分层进行。第一层是把整个 C 矩阵分成若干 M_BLOCK × N_BLOCK 的块,每个 block(线程块)负责一个输出分块;第二层是 block 内部把 K 维再切段,每次从显存加载小块 A、B 到共享内存;第三层是每个线程从共享内存取数据,用寄存器算对应的输出小片,累加结果最终写回显存。这个三层结构恰好对应 GPU 的三种存储层级:显存、共享内存、寄存器。

2.2 一个实例算清楚分块参数怎么定

举一个具体例子。假设目标 GPU 有 132 个 SM,目标形状是 M=N=K=4096。我最初选的 block 尺寸是 128×128,那 grid 就是 32×32 = 1024 个 block,平均每个 SM 要处理约 7.8 个 block,负载基本均衡。

K 维切片长度 BLOCK_K 的选择更讲究。BLOCK_K 越大,单次加载的数据越多,访问显存效率更高,但共享内存占用也随之增加。128×128 的 C 分块用 FP32 累加寄存器需要 128×128/32 线程 = 每线程 64 个 FP32 寄存器,加上操作数寄存器,寄存器压力已经不小。所以 BLOCK_K 我选了 32,这样 A 分块是 128×32×2 字节 = 8KB,B 分块是 32×128×2 字节 = 8KB,加上 C 分块,共享内存占用大概 20KB 左右,在 128KB 的共享内存里可以轻松放下双层缓冲。

提示:BLOCK_K 一旦超过 64,共享内存占用会急剧上升,留给双层缓冲的空间就紧张了。实际测试下来,在多数数据中心级 GPU 上,BLOCK_K 在 32 到 64 之间是最稳的选择区间。

选好这些参数之后,主循环的控制流就非常机械了:外层遍历 K 维切片,内层把切到的 A、B 分块从显存搬到共享内存,然后所有线程执行矩阵乘的小片计算。问题是,这种"搬一块、算一块"的做法搬数据的时候计算单元是空闲的,计算的时候搬运单元是空闲的,性能只有理论峰值的一半左右。解法就是第 3 节要说的双层缓冲。

3. 真正提速的核心细节:张量核指令、寄存器排布、流水线预取

3.1 mma 指令与 FP32 累加:为什么低精度输入要用高精度加法

现代数据中心级 GPU 架构都有一个专门做矩阵乘法的硬件单元,通常称为张量核心,它通过底层的 mma 指令一次性完成一个小矩阵的乘累加。比如一条 mma 指令可以完成 16×8×16 这样的操作,也就是 A 是 16×16、B 是 16×8,算出 16×8 的输出小矩阵。向量单元要一个时钟周期做几次乘加,张量核心一个周期能完成一个片段的乘累加,吞吐量完全不在一个量级。

DeepGEMM 里我采用的指令模式是:加载 FP16/BF16 输入,累加器用 FP32。这不是我拍脑袋定的——直接用 FP16 累加会在多次累加后出现明显的舍入误差,尤其在 K 很大时误差会累积。FP32 累加寄存器虽然占的地方多一倍,但换来的是结果精度明显提高,实测中与官方库的 FP32 累加结果完全相同。

PTX 层的指令看起来大概是这样:

mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32 {%0, %1, %2, %3}, // 输出累加器 16x8 的四个片段 {%4, %5}, // A 矩阵片段 16x16 {%6, %7}, // B 矩阵片段 16x8 {%0, %1, %2, %3}; // 输入累加器

这条指令的难点在于操作数布局是固定的,A、B、C 的每个片段分别放在哪些寄存器、哪些位,都有严格规定。如果寄存器排布不符合指令预期,编译器会补一堆 mov 指令来回倒腾数据,性能直接打折。所以我写 DeepGEMM 时是先按指令要求的寄存器布局来分配数据,而不是先分配好再期望编译器去适配。

3.2 共享内存 bank 冲突:看不见的串行化陷阱

共享内存按 bank 分块,硬件在同一周期可以同时服务多个不同 bank 的访问。但如果多个线程访问的是同一个 bank 的不同地址,吞吐就会变成原来的几分之一,这就是所谓 bank conflict。在 GEMM 内核里,加载 B 矩阵分块时最常踩这个坑。

经典案例:以 FP16 类型为例,BLOCK_N 为 128,B 分块是 BLOCK_K × 128。如果每个线程连续取一行中的相邻元素,两个线程的地址恰好落在同一个 bank 上,整个 load 就会被串行化。解决方法是给 B 矩阵的访问模式加一个偏置移位,让相邻线程访问的地址错开到不同 bank。

DeepGEMM 里我用了一个简单的 swizzle 策略:矩阵存储时,同一行的数据不再连续排列,而是按 XOR 变换重新组织。这样线程在列方向并行读取时,地址映射后的 bank 号自然错开,实测共享内存带宽利用率从约 60% 提升到接近 90%。

以 C 语言伪码来说明一个 8 元素 swizzle 的思路:

// 每个线程要读取的元素所在行 row、列 col // 传统布局:address = row * row_pitch + col // swizzle 布局:address = row * row_pitch \ // + ((col + 7) ^ ((row & 1) * 7))

关键是让"相邻线程的列地址"和"行本身的偏移"撞不到同一个 bank。

3.3 双层缓冲:让搬运和计算真正重叠

软件流水线的目标是让显存到共享内存的拷贝和矩阵乘计算重叠执行。最朴素的做法是共享内存里放两份 A、B 分块,一份给当前循环迭代使用,另一份给下一次迭代预取。

主循环写成这样:

而循环里的操作步骤是:

  1. 用 cp.async 指令异步发起下一次迭代的 A、B 数据拷贝
  2. 用当前共享内存缓冲区里的数据执行 mma 指令完成矩阵乘
  3. 提交计算完成等待之前的异步拷贝完成
  4. 交换两个缓冲区的角色,进入下一轮

cp.async 是 GPU 上的一种异步拷贝指令,数据从显存搬进共享内存不占用线程计算时间,硬件会自动完成。第一次实现时我直接把 cp.async 换成普通 load,性能下降接近两成,原因就是每次主循环结束都得等数据搬运完才能继续算。这个经验在 DeepGEMM 的开发中反复被验证:只要共享内存装得下,双缓冲几乎是零成本提升。

4. 精度对齐与排错链路:从错误结果到逐项排查

4.1 第一次跑通不代表结果是错的

我第一版 DeepGEMM 跑出了结果,乍一看和官方库的输出差不多,但差分对不过:最大绝对误差在 1e-1 量级,对推理模型来说这直接是错误输出。当时我最怀疑的是张量核指令用错了,后来发现其实是边界处理的问题。

排查的第一原则是固定变量。我把输入矩阵改成随机数,固定种子,先用官方库算出基准 C_ref,再用 DeepGEMM 算出 C_test,逐项对照。拿 M、N、K 都比较小的用例开始,比如 64×64×64,这样任何一行代码的行为都能手动推算。

4.2 排查流程:从索引、同步到边界

我当时整理的排查顺序,现在也推荐给你:

检查项方法我遇到的问题
索引公式在小矩阵上逐元素核对 A、B 分块加载确认 mma 指令的 A 矩阵行主序转列主序时搞混了一次
同步位置检查 __syncthreads() 是否覆盖跨线程共享数据读写有一次预取数据已经发起,当前计算还没用完旧缓冲区,就做了缓冲区交换
边界处理M、N 不能被 block 整除时,越界位置是否置 0128 的 block 处理 4096 没问题,换到 1000 就出错了
累加精度与 FP32 基准对比误差是否在可接受区间FP16 累加误差在 K=4096 时放大到不可接受

4.3 两个最容易踩的坑:越界加载与缓冲区交换

先说越界加载。BLOCK 尺寸通常是 32 或 64 的倍数,但矩阵 M、N 不一定是整数倍。当尾部块加载 A 分块时,最后几行已经超出矩阵边界,读到的显存内容是什么不确定,算出的结果自然不对。解法是加载时做边界判断,越界位置置零。这个判断放在主循环里太贵,我把它放在每次加载的 if 分支里:

if (row < M && col < K) { A_tile[local_row * BLOCK_K + local_col] = A[global_row * K + global_col]; } else { A_tile[local_row * BLOCK_K + local_col] = 0.f; // 越界置零,不影响累加结果 }

置零的好处是无论该位置被计算多少次,对最终累加结果都没有贡献,尾块照样可以走和完整块相同的 mma 指令。

再说缓冲区交换。双缓冲实现里有一个非常隐蔽的 bug:我在异步拷贝还没完成时就把缓冲区的指针换掉了,导致下次计算用的还是旧数据。排查这个问题的代码路径是用__pipeline_commit和__pipeline_wait_prior这对异步接口时,忘记在读取共享内存前等待对应批次全部完成。

注意:多级流水线里,等待条件不是"上一批拷贝完成",而是"我这次计算需要的那一批完成"。流水线越深,这个关系越容易搞错。

我后来用了一个比较实用的检验方法:把 K 维的切片数改成奇数跑一遍,如果结果和偶数切片不一致,基本就是流水线同步有 bug。因为切片数量奇偶变化会改变缓冲区和循环迭代的对应关系,任何等待顺序错误都会导致结果异常。

5. 性能实测:用 profiler 数据驱动,不拍脑袋调参

5.1 基线对比要分形状,不能只看一个数

GEMM 性能受形状影响极大。我拿 DeepGEMM 和官方闭源库做了对比测试,固定数据格式为 FP16 输入 + FP32 累加,分别测了三种典型形状,结果很能说明问题:

形状(M × N × K)DeepGEMM 达到的峰值算力占比与官方库相对性能
4096 × 4096 × 4096约 78%约 95%
1024 × 1024 × 4096约 70%约 88%
1 × 4096 × 4096约 15%约 20%

前两个结果说明 DeepGEMM 在常规大矩阵上已经有竞争力,最后一个 M=1 的形状则暴露了问题:block 里大量线程计算同一行,数据复用不够,内存加载变成了瓶颈。这个结果提醒我,DeepGEMM 目前的架构并不适合下沉到 M 很小的推理场景,还需要针对这种情况做独立优化。

5.2 用 profiler 找到瓶颈,而不是靠猜

性能出问题的时候,我先用 GPU 的性能分析工具采集三个指标:SM 忙碌率、共享内存带宽利用率、请求停滞分布。第一次跑 4096 形状,SM 忙碌率只有 55%,共享内存带宽利用率也不是特别高,说明问题不在计算而在流水线等待。

逐个排查后发现,主循环里每轮计算结束后有一个等待异步拷贝的操作,按我的预期这里应该已提前预取完毕。但分析工具显示 stall 主要发生在共享内存访问阶段,原因是加载 B 矩阵分块时的 swizzle 只做了部分偏移,仍有少量 bank 冲突。我把 swizzle 粒度从 4 元素调整到 8 元素之后,共享内存带宽利用率从 72% 上升到了 88%,整体算力占比也提升了大约 8 个百分点。

5.3 编译器参数同样影响性能

写 GEMM 内核时,编译参数的影响经常被低估。我用到的两个关键选项:一个是指定 GPU 架构让编译器针对具体指令集优化,另一个是开启更大的寄存器限制,允许寄存器溢出到局部的优化。

最开始我用默认编译参数,寄存器分配比较保守,性能差了约一成。把寄存器上限调高到 255 之后,编译器能把更多的中间结果留在寄存器里而不是反复写回共享内存。需要注意的是,这里有个平衡点:寄存器占用太高会导致 SM 上同时运行的线程块变少,并行度下降。我在目标 GPU 上实际测试,128 线程一 block、每个线程 168 个寄存器左右,是可以同时跑两个 block 的上限区间。

6. 往融合算子走一步:GEMM 的 epilogue 定制

6.1 推理场景真正需要的是 GEMM + 偏置 + 激活

如果 DeepGEMM 只能输出一个 C 矩阵,那它对推理引擎的价值就打折了。实际上全连接层之后的模式非常固定:先加 bias,再过 ReLU/GELU,最后才是输出。把这三个步骤从独立算子合并进 GEMM kernel 的末尾段(epilogue),可以减少一次完整的显存读写。

我在实现上把 epilogue 做成一个模板化的函数,传入输出块的共享内存指针和模型配置,由每个线程在算完自己的输出片段后执行。模板参数里包含是否需要 bias、使用哪个激活函数、是否要做量化,这些全都在编译期确定,运行时零分支开销。

6.2 量化场景:把 scale 和 zero point 焊进融合流程

低精度推理里,GEMM 输出是 FP32,但下一层要求 INT8 输入,中间必然有一次量化操作:y = round(x * scale + zero_point)。常规做法是把 C 矩阵写回显存,再启动一个量化 kernel 读出来重新算一遍,带宽和时间成本都很高。

在 DeepGEMM 的 epilogue 中,我在每个线程输出 FP32 片段之后直接做这个量化,再把结果写回显存。如果量化是 per-token 的,即每行一个 scale,也只需要把 scale 向量提前加载到共享内存,每个线程按自己的行索引取数即可。融合后一次 kernel 跑完 GEMM 和量化,实测整体的显存读写量下降约一半,端到端时间缩短了约 30%。

6.3 融合的边界与后续扩展方向

融合也不是越多越好。层归一化虽然也在 GEMM 之后常用,但它需要跨通道计算均值方差,意味着要拿到一行的完整输出,而 GEMM 的输出片段是按 block 分散在不同线程里的,强行融合会导致复杂的跨线程规约。我在 DeepGEMM 里没有把 LayerNorm 合进去,而是让它保持独立 kernel,这也是很多推理引擎的普遍做法。

这套内核后续可以扩展的方向,我目前最关心的是更小的 block 尺寸以适配推理批次较小的场景,以及把 FP8 输入的计算路径补全,因为在最新一代硬件上 FP8 的性能收益非常明显。另一个思路是为结构化稀疏做配套优化,稀疏矩阵在带宽上的节省潜力比纯稠密更大。

我实际写 DeepGEMM 的体会是,矩阵乘法的性能优化没有什么玄学,所有瓶颈最后都能落到访存模式、寄存器分配、流水线等待这几个具体问题上。关键是不要一开始就追求大而全,而是从一个小形状出发,把 profiler 给出的数据一项项磨平。你先跑通一个 block 都行,把同步、边界、精度都验证对了,再往上加复杂度和融合特性,这个过程会稳很多。

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

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

立即咨询