
很多人第一次看到DeepGEMM这个项目名第一反应是“又一个大模型框架”第二反应才是“GEMM 是什么”。这里先给一个最直接的回答GEMM 是 General Matrix Multiplication通用矩阵乘法的缩写也就是 ( C A \times B C ) 这类基础运算。深度学习里它有多重要全连接层、卷积层的 im2col 展开、注意力机制里的 QKV 投影本质上都是一次 GEMM。DeepGEMM 这一类项目的目标就是把矩阵乘法压到硬件极限在推理和训练时把延迟和带宽成本砍下来。这篇内容适合三类人刚接触高性能计算、想搞懂算子为什么能快几十倍的开发者已经在做推理优化、准备手写或魔改 GEMM 内核的工程向同学还有仅仅好奇“为什么同一个矩阵乘法有的库跑得快、有的库跑得慢”的算法从业者。我会从 GEMM 的瓶颈拆解开始讲到数据布局、切分策略、混合精度这些核心设计再给一套可以直接上手的内核实现思路和踩坑实录。整体偏实战会有不少伪代码和参数层面的讲解不涉及特定硬件平台的私有指令细节。1. DeepGEMM 到底在优化什么GEMM 的瓶颈拆解1.1 为什么深度学习里 GEMM 是绝对主战场先看一组很粗的账一个大语言模型的参数通常在几十亿到上百亿算一次 forward pass 要做多少次矩阵乘法以某个模拟项目 X 的 7B 模型为例单 token 推理时只算显式矩阵乘大概就有几十次 shape 在 ((batch, seq) \times (hidden, hidden)) 量级的主计算。训练场景下前向、反向、梯度更新各来一轮矩阵乘法占总计算量的比例能轻松超过 80%。可以说只要 GEMM 快 20%整个模型的吞吐就能肉眼可见地涨一截。这也是为什么几乎所有推理框架、训练框架里GEMM 相关算子都被单独拎出来做极致优化而不是像普通算子那样直接交给编译器默认生成。DeepGEMM 这种项目做的事情本质上就是那块最硬的骨头在给定的硬件上让矩阵乘法跑出接近理论峰值的实际性能。1.2 朴素矩阵乘法和极致实现差在哪里用最朴素的三重循环写一个矩阵乘法for (int i 0; i M; i) for (int j 0; j N; j) for (int k 0; k K; k) C[i][j] A[i][k] * B[k][j];逻辑上完全正确但性能往往只有理论峰值的 5% 到 10%。问题出在哪里两件事第一内存访问不连续内层循环里访问 B[k][j] 是在按列跳着读每次读一个浮点数就要等一次内存事务第二完全没有复用A[i][k] 被 C[i][j] 的每一列都读一遍B[k][j] 也被每一行读一遍数据在 CPU 缓存和内存之间来回搬运计算单元大部分时间都在等待数据。DeepGEMM 这类项目做的事就是围绕这两个问题做文章让访存变成连续的、块状的行为让数据在寄存器、L1、L2 缓存里尽可能重复利用最后才能把计算单元喂饱。理解了这个起点再看后面所有优化手段会发现它们全是在解决“数据喂不到计算单元嘴里”这个核心矛盾。1.3 计算峰值与访存带宽的账评估一个 GEMM 实现是否还有优化空间建议先做一次“纸面估算”。假设目标硬件上单精度浮点算力是 X TFLOPS当前矩阵乘法实际跑出来的时间是 T那么实际 FLOPs 就是 ( 2 \times M \times N \times K / T )。拿实际值除以理论峰值就是所谓的计算效率Compute Efficiency。但还有另一条线访存效率。一个 ( M N K 4096 ) 的 FP32 GEMM输入输出数据总量大约是[ 4096 \times 4096 \times 4 \times 3 201 \text{ MB} ]假如内存带宽是 200 GB/s光搬数据就要 1 秒而硬件算力如果支持在 0.1 秒内完成计算那就是典型的访存瓶颈。反过来如果数据能全部留在片上缓存里重复读取的次数就能从“每个输出都重新读一遍”降到“每个数据只读一次”这个差别在规模一大之后就是数量级的。所以每次拿到新的 GEMM 任务我第一件事永远是算这比账当前实现是计算受限还是访存受限再决定下一步优化方向。2. 设计思路与关键技术选型2.1 数据布局为什么行主序会“绊脚”深度学习框架里矩阵默认的行主序Row-Major布局看起来是合理的A[i][k] 和 B[k][j] 用同一个 i、k 索引访问逻辑上很好写。但到了高性能实现里这个布局会让 B 矩阵的内层循环直接变成跨行跳跃访问带宽利用率非常难看。所以 DeepGEMM 这类项目普遍会在 GEMM 内核里做一次布局转换B 矩阵从 ( K \times N ) 的行主序变成 ( N \times K ) 的列主序存储或者说等效于把 B 做了转置。这样内层循环访问 B 的一整行时地址是连续的。有人会问转置本身不也要花时间吗是的但转置可以放在初始化阶段或者算子融合阶段分摊到大量计算里代价几乎可以忽略。相比之下每次迭代都跳着访问内存的损失是要命级别的。实测下来在一个模拟的 FP32 GEMM 场景里仅把 B 转置成连续布局性能就能提升 20% 到 40%具体看矩阵规模。如果规模太小比如 ( 16 \times 16 )转置的启动开销会吃掉收益所以还得配合后面要讲的 Tile 切分来判断到底值不值。2.2 Tile 切分用尺寸换 Cache 命中和寄存器复用朴素三重循环的直接问题是内层循环的 k 维度跨度太远Scratchpad 和缓存根本接不住。标准解法是分块Tiling把矩阵切成一个个小的子块每个子块在计算时完全放进 L1/L2 缓存块内部反复读取时就不再触碰内存。先看一个简化但有效的分块策略。设分块尺寸为 ( BM \times BN \times BK )对应 A 的子块是 ( BM \times BK )B 的子块是 ( BK \times BN )。计算时外层循环遍历所有子块内层循环在一个子块内部做完整的 GEMM。常见的选择是 ( 64 \times 64 \times 16 ) 或者 ( 128 \times 128 \times 32 ) 这类尺寸。为什么选这些尺寸有一个经验公式单个线程/计算核的寄存器文件大小和 L1 缓存大小是硬限制。比如某目标平台有 256 KB L1那么 ( BM \times BK \times 4 BK \times BN \times 4 BM \times BN \times 4 ) 的总字节数必须远小于 L1。以 ( 64 \times 16 ) 的 A 子块和 ( 16 \times 64 ) 的 B 子块算总共是 ( 64 \times 16 \times 4 16 \times 64 \times 4 64 \times 64 \times 4 20480 ) 字节也就是 20 KB放进 L1 完全没问题。此时内部循环可以重复读这些数据很多次内存带宽的压力就大大缓解了。分块规模也不是越大越好。块太大缓存放不下会开始发生缓存驱逐反而退化成内存访问块太小复用的机会减少性能上不去。所以核心思路是让块的大小贴着一级缓存容量走同时确保每个块内的计算量足够大。2.3 混合精度与数值策略FP32 累加器的关键位置DeepGEMM 项目里另一个躲不开的话题是精度。现在 AI 推理普遍用 FP16 或 BF16 做输入输出如果直接拿两个 FP16 做乘加累加器还是 FP16很快就会出现精度灾难。标准解法是输入数据用 FP16/BF16 存储和读取乘出来的中间结果在寄存器里累加时使用 FP32 累加器。这样既节省了带宽和存储开销又保留了足够的数值动态范围。在具体实现里这一步通常意味着内层循环是这样的伪代码float acc[BM][BN] {0}; // FP32 累加器 for (int k0 0; k0 BK; k0) { half a_val A_block[i][k0]; half b_val B_block[k0][j]; acc[i][j] (float)a_val * (float)b_val; }关键点是“什么时候转成 FP32”。很多人在最初实现时图省事在读取 A、B 时就先转成 float再到寄存器里做乘加。这样做的坏处是内存读取带宽直接翻倍因为读进来的数据从 2 字节变成了 4 字节。正确做法是数据保持 FP16/BF16 读入在寄存器内部才展开成 FP32 做计算这样内存侧只承受 2 字节的带宽成本计算侧又有 FP32 的精度兜底。这一步对性能的影响很大在我实测过的一个模拟内核里单这个细节就能拉开 15% 以上的差距。2.4 为什么需要指令级手写而不是依赖编译器看到这里有人会问这些优化听起来都是 C 层面的思路编译器难道不能自动做吗答案是编译器能做一些循环展开和向量化但很难替你完成“寄存器级的分块决策”。因为 GEMM 的最优分块参数依赖于目标硬件的寄存器数量、缓存层级、SIMD 宽度这属于硬件特性的范畴编译器很难在通用层面生成最优解。所以 DeepGEMM 这类项目往往会在关键内层循环里使用 intrinsics——也就是直接调用硬件提供的向量指令或者在汇编层面手写一小段核心循环。这在普通应用代码里是大忌但在 GEMM 内核里是常规操作。因为内核代码只有几十行且在所有计算中重复执行的次数极高手写那几十行带来的收益远超维护成本。我个人的建议是第一步先用纯 C 把逻辑跑通、性能对标到“比朴素实现快 5 到 10 倍”的水平再考虑是否用 intrinsics 做最后 30% 的压榨。如果一上来就手写汇编调试难度会急剧上升而且很容易在没吃透 Cache 行为的情况下做出负优化。3. 实操过程手写一个可用的 DeepGEMM 核心计算内核3.1 从伪代码到可运行的 C 模板内核这部分给出一个可以直接参考的结构。假设目标环境是 x86 平台支持 AVX2数据类型是 FP16 输入、FP32 累加先写一个不依赖平台私有 intrinsics 的 C 版本重点讲清楚分块和布局转换怎么组织。第一步重排 B 矩阵布局。原始 B 是 ( K \times N ) 的行主序先把它转成 ( N \times K ) 的连续存储这样内层循环可以顺序读 B 的一行。// B_transposed[n][k] B[k][n] std::vectorhalf B_T(N * K); for (int k 0; k K; k) for (int n 0; n N; n) B_T[n * K k] B[k * N n];第二步按 Tile 遍历。外层循环处理所有子块内层把子块数据搬进局部缓冲区然后再做局部 GEMM。先不用管复杂性把逻辑写清楚再慢慢优化。constexpr int BM 64, BN 64, BK 16; for (int i0 0; i0 M; i0 BM) { for (int j0 0; j0 N; j0 BN) { float acc[BM][BN] {0.f}; for (int k0 0; k0 K; k0 BK) { // 把子块 A[i0 : i0BM][k0 : k0BK] // 子块 B_T[j0 : j0BN][k0 : k0BK] 载入局部内存 half a_block[BM][BK]; half b_block[BN][BK]; // ... load with boundary checks ... for (int i 0; i BM; i) for (int j 0; j BN; j) for (int k 0; k BK; k) acc[i][j] (float)a_block[i][k] * (float)b_block[j][k]; } // 写回 C[i0 : i0BM][j0 : j0BN] } }核心逻辑已经在这里了。接下来要做的所有优化都围绕“怎么让内层那个三重循环跑得更快”把 a_block 和 b_block 的读取改成连续把 acc 数组尽量放到寄存器变量里把循环展开到合适的粒度。一个重要的实操点不要一上来就写边界检查因为每次内层循环里加 if 判断会有分支开销。正确做法是把矩阵填充到对齐的尺寸比如 M 向上取整到 BM 的整数倍填充部分用 0这样内层循环就是完全无分支的。这也是我反复强调的“用内存换分支”思路。3.2 极端 Case 处理与 Shape 泛化实际业务里不会总是方方正正的矩阵M、N、K 可能是任意整数。所以 DeepGEMM 的实现必须处理“边界”问题。处理方式有两种主流方案第一种填充法上面已经提到。把 M 向上对齐到 BMN 向上对齐到 BNK 向上对齐到 BK。新增的部分用 0 填充不影响计算正确性。这个方案实现最简单缺点是一次性分配多出一点内存。对于推理场景中的固定 shape这是最优解因为填充只在启动时做一次。第二种分支法在每次加载子块时判断边界里的有效数据数量。这个方案对任意 shape 都成立但内层循环里有条件跳转性能会略有下降。实测下来当 M、N 恰好是 Tile 整数倍时分支法的性能可以做到和填充法几乎一致但如果不是分支法可能会慢 10% 到 20%。我的建议是框架内部统一用填充法因为算子层的 shape 基本是编译器静态已知的。如果写的是动态 shape 的通用算子就在 dispatch 层套一个分支shape 对齐时走快速路径不对齐时走通用路径两条路径共用同一套内核模板。此外还有 A 矩阵和 B 矩阵的转置标志问题。深度学习里经常出现“对 A 做转置”或者“对 B 做转置”的需求比如 attention 里的 score 计算就常常是 ( Q \times K^T )。如果每种布局组合都写一个内核代码量会爆炸。标准做法是统一把输入布局在初始化阶段重排成内核需要的最优布局内核本身只支持一种布局。这个思路看起来多了一步拷贝但胜在 Debug 简单而且可以规避大量分支。3.3 指令级优化方向SIMD 与寄存器分块当 C 版本的逻辑稳定后就可以开始考虑指令级优化。这一步的核心目标是把内层循环变成每次从内存取一批 A 子块和 B 子块数据在寄存器里做一批乘加再写回。我用 AVX2 平台举例思路在其他 SIMD 平台通用。AVX2 的向量宽度是 256 bit可以一次处理 8 个 FP32或者 16 个 FP16。如果做 FP16 输入、FP32 累加一种方案是把 16 个 FP16 拆成两组 8 个转成 FP32 后分别用两条向量指令做乘加。伪代码大致是__m256 acc_vec[BN_VEC]; for (int k 0; k BK; k 8) { __m128i a_half _mm_loadu_si128(...); // 8个FP16 __m256 a_float _mm256_cvtph_ps(a_half); // 转成8个FP32 // 对 B 做同样操作 for (int jj 0; jj BN_VEC; jj) { acc_vec[jj] _mm256_fmadd_ps(a_float, b_float[jj], acc_vec[jj]); } }这个方向优化的空间非常大核心在寄存器分块参数的选取。寄存器数量是有限的你要决定一次循环里同时算多少行、多少列的输出。比如把 ( 8 \times 8 ) 的 C 子块放在寄存器里那么需要 8 个向量寄存器放累加器另外还要几个寄存器放 A 片段和 B 片段。如果贪心选 ( 16 \times 16 )寄存器不够用编译器就会把部分累加器溢写到栈上性能反而暴跌。我踩过的坑是一开始为了“看起来更并行”选了很大的寄存器分块结果寄存器溢出导致性能比编译器自动生成还差。调试了很久才发现加一个-fverbose-asm看一下汇编里面全是vmovups和栈操作根本不是计算指令。后来老老实实从 ( 8 \times 8 ) 开始测试再逐步加大才找到了当前平台的最优参数。4. 性能验证与常见问题排查4.1 测试工具与性能基线对比方法评估一个 GEMM 内核是否合格不能只看“跑得快”还要看相对基线的提升幅度。我建议准备三组数据做对比第一组是朴素三重循环实现直接作为性能下限基准。第二组是当前框架里成熟的 BLAS 库实现作为业界水准参照。第三组就是你自己写的 DeepGEMM 内核。每次修改只动一个小变量保持其他条件不变记录时间。性能测试的步骤也很关键。矩阵乘法的耗时受数据冷热状态影响极大第一次调用会把数据从主存搬到缓存后续调用会命中缓存时间明显变短。正确做法是先跑若干次 warmup让缓存状态稳定然后跑 50 到 100 次计时取中位数而不是取最小值或平均值。平均值会被偶发的系统调度干扰最小值往往只是冷启动时碰巧没被中断打断中位数最稳定。FLOPs 的计算也要准确。一次标准的 ( M \times N \times K ) GEMM浮点操作次数是 ( 2 \times M \times N \times K )乘法和加法各一次。用计时时间去除这个数就得到实际算力。注意很多人在 K 维度上算错把 ( 2 \times M \times N \times (K-1) ) 或漏掉加法操作导致百分比虚高。4.2 排查实录一数值不一致越优化越错一个常见的诡异问题是内核写对了跑小矩阵时结果正确但矩阵一大结果就开始出现微小的偏差有时甚至直接出现 NaN。这类问题第一嫌疑是边界越界。当 M 不是 BM 的整数倍而填充又没做好时内层循环会读到未初始化内存可能是任意值包括极小的浮点数或规格化数累积到 FP32 累加器里就会产生不可控的误差。排查方法不难在 main 函数里把整个输入矩阵初始化为 0并开启 AddressSanitizer 跑一遍如果提示堆缓冲区溢出基本就是这里。第二嫌疑是 FP16 表达的动态范围不够。某些模型权重绝对值在 ( 1e-4 ) 到 ( 1e-4 ) 之间的稀疏值很多转成 FP16 后直接变成 0或者产生非规格化数累加后和原始 FP32 结果差得多。这种情况不能怪内核只能改用 BF16 或者在算子层面做误差补偿。第三嫌疑是 cache line 伪共享多见于多线程版本里多个线程写同一个缓存行。比如两个线程各自算 C 矩阵相邻的两列如果列间距很小它们会频繁争抢同一个 64 字节缓存行速度急剧下降数值也可能因为非原子操作而互相覆盖。解决思路是每个线程负责的 C 子块至少按缓存行对齐比如每个线程至少处理 8 个 FP32 列避免共享。4.3 排查实录二性能倒挂手写内核反而更慢很多人第一次把 DeepGEMM 写完后去对标 BLAS 库结果发现自己的版本慢得离谱甚至比朴素实现还慢。这里我梳理几个高频原因第一个原因是内存分配的开销。每个子块都new一块局部内存循环里反复分配释放性能基本就废了。正确做法是在内核外部一次性分配好所有工作缓冲区复用同一块内存。第二个原因是编译器优化没有生效。C 模板内核看起来逻辑正确但编译器可能没有把局部数组放到寄存器里而是一股脑地放到了栈上。用-O3 -marchnative编译是基本操作另外可以试着把内层循环里访问的数组声明为register或使用局部变量让编译器更容易识别。更保险的做法是检查生成的汇编里内层循环有没有向量化指令如果全是标量操作就说明编译器根本没听你的。第三个原因是你可能没有启用多线程。现代 CPU 的 GEMM 性能峰值是多核协同的结果单核实现即便做得再极致也顶多达到一部分性能。但多线程不是简单加个#pragma omp parallel for就完事的线程数、任务划分粒度、线程亲和性都影响结果。我通常的实践是先验证单核版本已经达到单核理论峰值的 70% 以上再开始做多线程扩展否则问题叠加起来很难排查。第四个原因是矩阵规模太小。GEMM 内核的启动开销和 Tile 切分开销是固定的当 M、N、K 只有几十的时候这些开销会淹没计算收益。这类规模在业务里也很常见比如某些很小的全连接层。针对小矩阵另有专门的 kernel 设计策略通常不会走BM64这种大切分路线。4.4 避坑清单DeepGEMM 开发中的高频问题和应对整理一份简表按影响程度排列问题主要症状应对方案B 矩阵布局未重排性能远低于预期统一转换为 ( N \times K ) 连续布局内层循环有边界 if性能不稳定、分支预测开销高填充到对齐尺寸消除分支局部数组导致栈溢出编译后的汇编全是 vmovups使用寄存器变量或手动展开多线程伪共享核数增加但性能不升反降按缓存行对齐任务划分FP16 累加精度不足大矩阵结果偏差明显切换 BF16 或使用 FP32 累加器冷缓存计时性能数据波动大warmup 后取中位数寄存器分块过大汇编出现栈溢出操作缩小寄存器分块逐步逼近最优这份清单基本覆盖了我自己在做 DeepGEMM 类项目时遇到的绝大多数问题。有些问题光看代码很难发现但只要对汇编和内存行为有一定敏感度可以靠经验快速定位。5. 从算子到系统的扩展方向5.1 融合算子与 DequantGEMM 绝不单打独斗GEMM 内核本身的优化做到一定程度后边际收益会越来越小。这时候真正的系统级优化开始登场把 GEMM 前后的算子融合进同一个内核减少中间结果的写回和读入。典型例子是量化模型的 Dequant 融合如果权重是 INT8 或 FP16而计算需要反量化回高精度传统做法是先跑一遍反量化算子把结果写到内存再跑 GEMM。融合方案则是在加载 B 子块数据时直接在寄存器里做反量化然后马上参与乘加。这个操作省掉的是一整轮 ( K \times N ) 矩阵的读写收益极其可观。融合后的实现需要注意寄存器压力。Dequant 操作通常需要额外的查找表或缩放因子处理不好就会挤占累加器的寄存器空间。实践上可以多测几种融合策略比如有的平台把缩放因子做成标量乘法有的平台要做向量查表选在当前硬件上表现最稳的那种。5.2 多核并行与负载均衡当单核内核稳定后下一步就是让多个计算核同时干活。任务划分的核心原则是每个核处理一块尽量独立的 C 子块避免共享写入。常见做法是把 C 矩阵按行方向切分成若干连续带每个核负责一条带。如果矩阵高度 M 非常大可以做两层切分第一层按行切分出大任务块第二层在块内再按列切成小片便于动态负载均衡。负载不均衡的问题在真实场景里很常见。比如 M100、N10000 这种形状如果按行切分某些行因为没有足够 K 维度计算耗时很短某些行很长导致一部分核提前空闲。应对方案是使用任务队列而不是静态切分每个核完成一个任务片后去队里取下一个这样就算任务大小差异很大总体也能保持较好的负载均衡。我自己的经验是做多核扩展时先在单核版本上用perf统计出计算和访存的比例确认瓶颈类型。是访存受限时加再多核收益也有限因为内存带宽已经被吃满了是计算受限时加核心数才有效果。很多项目组调了半天线程数没有提升其实是没意识到早就是带宽瓶颈了。5.3 值得继续深挖的几个延伸方向DeepGEMM 这类项目可以延伸的方向其实不少简单列几个我觉得后续价值比较高的稀疏 GEMM。大语言模型里越来越多的权重是稀疏的把零值跳过不计算能进一步压计算量和带宽。但稀疏矩阵的存储格式和并行策略跟稠密版本差异很大建议单独开一个子项目去探索。单算子多版本自动调优。不同 shape 下最优 Tile 参数完全不同可以把常见 shape 对应的最优参数离线搜索一次存成配置表运行时按 shape 查表选择版本。与编译器协同。手写内核虽然性能极强但开发成本高。如果能把你验证过的最优 Tile 参数和布局策略沉淀成编译器 hint让编译器在更广泛的算子上自动生成类似代码对整个框架的性能都是正收益。这些方向里自动调优是性价比最高的。因为 GEMM 的优化空间是分层的只要把某一层的最优参数固定下来后续的新硬件支持就只需要重新跑一遍搜索流程不用再重新设计内核。回到开头那句话GEMM 是深度学习的绝对主战场。DeepGEMM 这类项目教会我的最重要一课是性能优化不是玄学而是“理解硬件 拆解访存 反复实测”的工程循环。如果你正准备开始做自己的 GEMM 内核我的建议是别一上来就追求极致先把 C 版本跑通把性能对齐到朴素实现的 5 倍以上再来触碰指令级优化——等你走到那一步会发现之前所有在布局和切分上花的时间全部会折算成最终性能的一部分。最后分享一个实用技巧调试 GEMM 内核时永远用一个很小的例子比如 8×8×8和一份朴素实现做对拍每次改完代码都先跑这个小例子的结果对比再跑大规模性能测试。这一步能防止你把性能优化和正确性调试搅在一起省下的时间远比你想象的多。