DeepGEMM深度解析:从GEMM底层原理到Tensor Core极致性能优化
矩阵乘法GEMM这个话题在深度学习圈子里这几年越来越热。最近看到不少人在讨论 DeepGEMM 这个名字乍一听像是给 GEMM 套了个“深度”的前缀实际上它指代的是一类专门为深度网络推理与训练场景设计的高性能矩阵乘法方案。我可以直接说结论深度学习框架里跑得最多的底层运算就是 GEMM大模型里的注意力、MLP 投影、卷积的 im2col 变换本质上全是在做矩阵乘法。通用 BLAS 库虽然很强但面对深度学习这种形状固定、精度要求特殊的场景反而会有不少浪费。DeepGEMM 这类项目解决的就是“如何在特定形状、特定精度、特定硬件上把矩阵乘法压到极限”的问题——它不仅适合系统工程师和研究性能优化的人看也适合所有想搞懂“Transformer 为什么快/慢”的算法工程师。我打算用一篇完整的实践拆解来聊这个主题从 GEMM 的底层原理讲起到分块、流水线、Tensor Core 调用方式再到调试排坑和性能分析。内容里会带上可落地的代码骨架和调参思路而不是停留在概念层面。1. 整体设计与思路拆解深度学习专用 GEMM 到底在优化什么1.1 先走一遍 GEMM 的物理本质任何矩阵乘法都长这样C[M, N] A[M, K] × B[K, N]顺便可以加偏置。M 通常表示样本或 Token 数N 表示输出特征宽度K 表示归约维度。Transformer 里 attention 的 QKV 投影是 GEMMMLP 的两层全连接是 GEMM位置编码的旋转矩阵拼接也可以看成 GEMM。所以 GEMM 占大模型训练和推理总计算量的 90% 以上一点不夸张。但是 GEMM 有一个物理瓶颈叫做“访存墙”。CPU 和 GPU 的算力增长速度远快于内存带宽导致很多核函数并不是“算不过来”而是“数据搬不过来”。这就是为什么 GEMM 优化的核心不是算术本身而是怎么让数据在正确的时间出现在正确的位置。一个标准矩阵乘法需要读 A 和 B写 C。如果 M、N、K 都是 4096FP16 精度下总数据量约为 100MB 的量级GPU 显存带宽按 2TB/s 算光搬数据就要几十微秒。而 Tensor Core 算完这些乘加只需要几微秒差了一个数量级所以不搞优化根本跑不满硬件。DeepGEMM 这个方向的核心思路就是在“全连接形状固定、精度固定、硬件固定”的前提下把 GEMM 的每一份访存都用到极致。通用 BLAS 库为了兼容各种形状、各种布局、各种 stride内部会做很多判断和分支这在深度学习场景下反而是负担。你明明知道 B 矩阵一定是 K×N 连续排布还去处理转置和跨步纯属浪费。1.2 为什么通用库不是万能的DeepGEMM 的价值在哪很早以前大家写神经网络底层调用都是直接用厂商闭源的 BLAS 库。闭源库确实强但有几个问题让人头疼。第一个问题是形状适配。厂商 BLAS 库要覆盖从 1×1 到大到离谱的矩阵内部采用启发式选择算法。但在深度学习中很多场景 M 是固定的比如 batch 是 32 的时候 M32推理场景 M 就是 1 或者 8。这种“尾巴形状”在通用库里往往不是最优路径性能可能比预期掉 20%-50%。第二个问题是指令集适配。Tensor Core 系列的硬件指令更新换代很快闭源库在上一代架构上可能优化得很好到新一代硬件上初期就跟不上。手写 GEMM 内核对厂商的指令集更新可以做到“发布会当天就跟上”社区项目经常有这样的操作。第三个问题是精度策略。推理场景流行 FP8、INT8 量化训练场景流行 TF32、混合精度。通用库给的 API 大多是某一套固定策略DeepGEMM 这类项目则可以把量化缩放、误差补偿、尾数截断全写进内核里和上层的量化算子深度配合。所以 DeepGEMM 不是“造轮子”而是“定制轮子”针对一组有限场景做极致的定制牺牲通用性换性能。这是我个人认为这类项目最有价值的地方。如果你只在 CPU 上跑小型矩阵别碰这个方向如果你需要把某几组固定形状的 GEMM 跑到硬件极限方向就对了。2. 核心细节解析与实操要点分块、复用、流水线三条主线2.1 分块为什么不能直接算完整个矩阵GPU 里有很多计算单元每个计算单元访问的是一小块片上的高速缓存。如果每个线程都要读整个 A 和整个 B那数据从显存搬进缓存几百次都搬不完。分块Tiling的思路就是把大 GEMM 切成很多个小块每个块只负责其中一小段计算让每个小块都能完整装进片上缓存。我拿“图书馆搬书”来类比。你要做一个很大的计算题需要参考很多本书。如果每算一步就跑去楼下图书馆翻一本书速度极慢。分块优化就是先把书架上某个区域的几十本书全部搬到你的桌子上然后在这个桌上连续工作很久等这批书用完了再去搬下一批。GPU 里的“桌子上”就是共享内存。具体来说假设我们有 A 的 128×32 一块和 B 的 32×128 一块这两块做完乘加可以得到 C 的一个 128×128 分块。32 是归约维度的块大小记作 BK。为什么取 128 和 32不是拍脑袋。共享内存容量是有限的一般可配置到 100KB 以上。FP16 下 A 块大小 128×32×2字节 8KBB 块 32×128×2字节 8KB两块加起来 16KB。如果做双缓冲要 32KB剩余空间还能放中间结果完全容纳得下。同时 128×128 的分块规模能产生足够多的独立小块让所有计算单元忙起来不至于算完一块之后大家闲着等数据。块尺寸也不是越大越好。取 256×256 的话单个块光共享内存就要 256KB直接爆了只能去分成子块反而增加复杂度。取 64×64 的话共享内存确实随便用但计算单元之间的数据复用率降低从全局内存加载数据的次数变多带宽压力变大。我在实际测试中128×128 或 128×64 是大多数核函数首选的起点后面再按硬件微调。2.2 数据复用与累加寄存器GEMM 还有一个容易被新手忽略的点C 矩阵的每个元素需要累加 K 次。如果每次累加都去显存里读一下旧值再写回整个带宽直接爆炸。所以正确做法是在寄存器里维护一段 C 分块让所有累加都发生在寄存器层面。寄存器是每个线程私有的速度最快容量也最小。一个线程通常维护一个 8×8 或 16×8 的微块tile也就是要开对应数量的浮点寄存器来保存累加值。举个例子一个线程维护 8×8 的微块就需要 64 个浮点寄存器。加上其他临时变量线程占用在 128-200 个寄存器之间。这个数字不能太高否则一个 SM流式多处理器上能同时运行的线程数下降调度灵活性降低反而影响隐藏访存延迟的能力。分块、寄存器累加、向量化加载这三点构成了所有高性能 GEMM 的公共底座。在此基础上各家方案再往数据流水线上堆叠技术。我下面会展开讲流水线。2.3 计算与访存重叠的流水线设计刚才说的“屋子里搬书”还有一个进阶版本不等用完一批书才去搬下一批而是在用当前这批书算题的同时安排人先去搬下一批书。这就是双缓冲double buffering也叫流水线。在 GPU 里显存数据搬到共享内存走的是异步拷贝指令。理想的状态是计算单元执行当前块的乘加指令时片上缓存织入指令已经在后台把下一块所需的数据从显存搬到共享内存的另一个缓冲。计算和访存完全不互相等待。DeepGEMM 这类项目通常把流水线分为几个阶段读全局内存global load、写共享内存shared store、计算乘加mma、读下一块prefetch。在实践中我会把主循环写成“提前加载 循环计算 尾部处理”的结构也就是说主循环开始前先发出第一块数据的异步加载然后主循环里依次执行“计算当前缓冲 → 加载下一块缓冲 → 同步屏障”这样每次迭代的加载操作都能和上一轮的计算重叠。流水线段数也可以从两级加深到三级甚至四级。不过级别越多共享内存占用越高、同步控制越复杂。两条缓冲在绝大多数场景下性价比最高三条以上只在极端的带宽模型下有意义。我试过把一个核函数的缓冲从 2 加到 3性能提升不到 5%但调试成本明显上升所以默认建议 2。2.4 精度模式选择从 FP16 到 TF32再到 FP8深度学习专用 GEMM 的另一个看家本领是精度策略。很多人以为“混合精度”就是 FP16 存、FP32 累加实际上这只是最基础的方案。FP16 的问题在于尾数只有 10 位。当累加值较大、乘积累加数值较小时精度容易被吃掉。常见对策是 K 维分片累加每片独立累加后用 FP32 累加器做合并。另一种是“3x FP16”技巧把 FP32 拆成高、低两个 FP16 分别表示三个乘法组合逼近 FP32 的精度这个在 DeepGEMM 类项目里很常见。TF32 是某厂商提出的截断格式本质上是 FP32 的指数范围加上大约 19 位有效精度截掉了尾数的一部分用 Tensor Core 跑比纯 FP32 快很多。代价是精度下降了所以训练场景一般把误差敏感层留在 FP32其余层用 TF32。FP8 则更进一步它有 E4M3 和 E5M2 两种变体。前者精度更高适合前向后者动态范围更大适合反向。FP8 GEMM 通常要配合 per-tensor 或 per-token 的缩放因子把数值范围压缩到可表示区间。这已经不是单纯换数据类型了而是要在内核里做 rescale。DeepGEMM 项目里这部分往往写成模板参数让调用方选择是否做缩放以及缩放模式。精度模式有效位宽累加方式速度相对FP32适用场景FP3223位尾数FP321x精度敏感层TF32约19位FP324-8x训练前向/反向FP1610位尾数FP328-16x常规训练FP8 E4M33位尾数4位指数FP32需缩放16-32x推理、量化训练FP8 E5M22位尾数5位指数FP32需缩放16-32x梯度反传需要提醒的是精度模式的选择永远要跟业务效果挂钩。我见过有人把 FP8 GEMM 直接应用到某个对数值范围极敏感的模型上收敛崩了随后反过来骂框架。其实问题出在缩放因子的更新策略上和 GEMM 本身的实现关系不大。这一点在后面的“数值误差”部分会继续展开。3. 实操过程与核心环节实现从朴素核函数到一个能跑的优化版本3.1 明确目标形状和硬件约束动手写 GEMM 内核之前先定义一个具体场景。我拿最常见的推理场景举例M64N4096K4096FP16 输入FP32 累加。这个场景在视觉模型的检测头、大模型的输出投影层里非常典型M 是 batch×beam 或者 token 数的小批。然后看硬件能力。假设一块主流商用 GPU支持 Tensor Core 的 mma 指令共享内存可配置 96KB线程束大小 32。在这个前提下目标很明确让 Tensor Core 干活不让它等数据。也就是做到计算密度足够高、访存延迟被隐藏。为了快速验证思路建议先用高层内核语言写一个原型像 Triton 或者类似工具在原型上调通逻辑确认性能量级后再用底层语言去抠细节。这是 DeepGEMM 类项目最常见的推进路径能省大量时间。3.2 原型实现先让正确性跑通下面是一个伪代码级别的原型。我先忽略 swizzle 和异步拷贝先把计算结构摆出来。import triton import triton.language as tl triton.jit def deepgemm_kernel( A_ptr, B_ptr, C_ptr, M, N, K, stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_cn, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, ): pid_m tl.program_id(0) pid_n tl.program_id(1) offs_m pid_m * BLOCK_M tl.arange(0, BLOCK_M) offs_n pid_n * BLOCK_N tl.arange(0, BLOCK_N) offs_k tl.arange(0, BLOCK_K) a_ptrs A_ptr offs_m[:, None] * stride_am offs_k[None, :] * stride_ak b_ptrs B_ptr offs_k[:, None] * stride_bk offs_n[None, :] * stride_bn acc tl.zeros((BLOCK_M, BLOCK_N), dtypetl.float32) for k in range(0, tl.cdiv(K, BLOCK_K)): a tl.load(a_ptrs) b tl.load(b_ptrs) acc tl.dot(a, b, out_dtypetl.float32) a_ptrs BLOCK_K * stride_ak b_ptrs BLOCK_K * stride_bk c_ptrs C_ptr offs_m[:, None] * stride_cm offs_n[None, :] * stride_cn tl.store(c_ptrs, acc)这段代码的思想每个 program 负责一个 BLOCK_M×BLOCK_N 的输出块在 K 维度上循环不断加载 A 块和 B 块用 tl.dot 做矩阵乘累加。BLOCK_M128、BLOCK_N128、BLOCK_K32 的时候编译后会自动去匹配 Tensor Core 指令。原型跑出来的性能通常已经超过普通的逐元素实现但距离手写极限还有一段距离。跑通正确性之后再看 profiling 结果。如果加载带宽不满、计算利用率低再进入手写内核的流程。3.3 手写内核分块、双缓冲与异步拷贝底层手写时我会按下面这个骨架来组织代码。先把主循环的流水线结构摆清楚。// 伪代码展示 DeepGEMM 类内核主循环的核心结构 // CTA 负责计算 block_C 128x128 // block_A: 128x32, block_B: 32x128, FP16 __shared__ half smem_a[2][128][32]; __shared__ half smem_b[2][32][128]; // 第一步预加载第0块到 buffer 0 load_to_smem(0, threads, smem_a[0], smem_b[0]); for (int k 0; k K / 32; k) { // 当前计算使用 buffer k % 2 // 同时异步加载下一块到 buffer (k1) % 2 if (k 1 K / 32) { load_to_smem_async((k 1) % 2, ...); } __syncthreads(); // 每个线程基于当前共享内存块调用 mma 指令 for (int tile_m 0; tile_m 4; tile_m) { for (int tile_n 0; tile_n 4; tile_n) { mma_sync(acc[tile_m][tile_n], smem_a[k % 2][tile_m * 32][...], smem_b[k % 2][...][tile_n * 32], acc[tile_m][tile_n]); } } __syncthreads(); // 确保计算完成buffer 才能被覆盖 }这个骨架有三个关键点。第一异步加载指令不能和普通赋值混在一起需要用底层拷贝指令把数据从全局显存搬到共享内存。代码层面要确保加载和计算使用的是不同 buffer才能实现重叠。第二每个线程持有的累加寄存器和 mma 指令的排布要一一对应。一个 warp 32 个线程一起执行 mma 指令拿 8×8 的微块来说32 个线程各持有若干寄存器片段不能随意倒腾位置否则跨线程的数据交换会非常昂贵。第三两个 __syncthreads() 一个都不能少。第一个保证共享内存中的数据已经就绪第二个保证所有线程都算完了再覆盖缓冲。漏掉一个轻则性能下降重则数据竞争导致随机错误。这类问题在 profiling 里不太好查多花点心思在同步上。3.4 进一步优化Swizzle 和 Bank Conflict到了这一步大部分计算资源已经被利用起来了剩下的收益通常在共享内存的访问冲突上。这就是常说的 bank conflict 问题。共享内存被划分为若干个 bank同一时钟周期内如果多个线程访问同一个 bank硬件就会串行化处理性能直接打折。避免 bank conflict 的常见做法是给共享内存的下标做 swizzle 变换。简单说就是让 A 矩阵分块在共享内存中的放置方式不那么“连续”而是按某种系数打乱。比如smem_a[row][(col ^ row) % 32]这种写法的变体可以让同一行线程访问的地址落到不同 bank 上。这个细节优化在某些硬件上能带来 10%-20% 的差距但实现时对下标计算的额外开销也要控制好不能“优化了个寂寞”。实测下来对于维度比较大的矩阵把共享内存的数据布局从普通的行主序改成 padded填充模式就够了。每行多填充几个 half让行的起始地址错开往往能消除大部分 bank conflict。padded 模式实现简单、不易出错我建议新手先从这个入手。3.5 调参清单与实际效果对比给一个我常用的调参顺序每一步都有明确的检查指标。第一步选定 BLOCK_M 和 BLOCK_N 为 128。看 SM 占用率是否达到 80% 以上。如果达不到检查是不是寄存器数量爆了适当降低每个线程维护的微块尺寸。第二步选定 BLOCK_K 为 32 或 64。BLOCK_K 越大数据复用率越高但共享内存压力也越大。对于 FP1632 是比较稳的起点64 需要硬件共享内存足够大。第三步调整线程束数量。一个 128×128 的输出块如果用 8×8 微块需要 16×16256 个线程块也就是 8 个 warp。如果要提高每个线程的计算量可以改到 4 个 warp每个线程维护更大的微块。一般来说 4 个 warp 的同步开销更小但寄存器压力更大。第四步打开 profiler 看数据。重点看三个指标SM 占用率、Tensor Core 利用率、显存带宽利用率。性能优化不到位的 GEMM往往能在这三个指标上直接看出短板。不要凭感觉调要用数据说话。我给一个参考性能变化趋势一个朴素版本无分块、无共享内存在 FP16 大矩阵下可能只能跑到几个 TFLOPS加上分块和共享内存后能到数十 TFLOPS再加上双缓冲和 swizzle 后接近硬件的峰值比例。注意这只是量级参考不同卡差异很大。4. 常见问题与排查技巧实录我踩过的那些坑4.1 性能上不去先别怀疑代码先看瓶颈在哪类我第一次优化 GEMM 内核时花了很久在各种微优化上结果性能纹丝不动。后来拿 profiler 一看问题根本不在计算端而在全局内存加载带宽已经被占满了。也就是说算法已经把 A、B 数据重复读取了很多次而共享内存复用率过低。方向错了再怎么优化 microbenchmark 都没用。我建议拿到任何 GEMM 优化任务第一件事是算一下算术强度arithmetic intensityFLOPs 除以涉及的字节数。千万别跳过这个步骤它能直接告诉你是访存密集还是计算密集。深度学习 GEMM 大部分是访存密集的所以优化方向要优先落在减少重复加载上。计算方法是C 矩阵 M×N×2字节A 矩阵 M×K×2字节B 矩阵 K×N×2字节。假设精度 FP16总访存量就是 2×(M×N M×K K×N) 字节。总计算量是 2×M×N×K FLOPs。二者相除得到算术强度。如果低于硬件 FLOPs/带宽的比值那必然是访存瓶颈。这个计算会直接影响你选多大的块算出来访存密集就把块尽量调大算出来计算密集就聚焦 Tensor Core 利用率和指令调度。4.2 数值误差比预期大的排查思路有段时间我在调一个混合精度 GEMMFP16 输入、FP32 累加结果误差一直超标。查了很久发现问题出在 K 维的累加顺序上。不同的 accumulate 顺序产生不同的舍入误差特别是 K 很大时误差会叠加。解决办法很简单把 K 维循环拆成多段每段累加一个局部结果最后再合并局部结果。这在代码里就是多开一组中间寄存器在尾部做一次归约。代价是额外占用寄存器但精度提升明显。还有一次FP8 场景误差大得离谱。检查发现缩放因子被在线性层之前施加了两次等于把输入数值范围压得过小精度全部丢失。排查这类问题时建议写一个很小的单元测试固定输入对比 FP32 参考输出和优化输出的误差分布。不要只比较最大值要比较均方误差和逐元素误差。有时候最大误差来自个别极端数值均值误差反而没问题。4.3 显存访问冲突导致的隐形降速共享内存 bank conflict 不像 bug 那样会报错它只是默默拖慢速度。最典型的特征是profiler 显示共享内存吞吐接近 100%但计算单元利用率很低。这个时候第一反应就应该是冲突。我排查 bank conflict 的快捷方法先用 profiler 看 shared memory 的相关计数器如果数值异常就去审查共享内存数组下标的索引表达式。最简单有效的修复是给每个行增加 padding。对于 FP16 的 128×32 数组可以把每行宽度从 32 改成 34 或者 36。这样下一行的起始地址偏移了一个 bank大量冲突就消失了代价只是多占一点点共享内存。更高级的 swizzle 方法我前面提过这里不重复。想说的是不要一开始就上 swizzle先用 padding 拿到收益确保正确再根据 profiling 逐步替换。4.4 常见问题速查表现象可能原因优先排查方向计算利用率低Tensor Core 没被调用或调用次数太少确认是否真正发出 mma 指令而不是标量乘加带宽利用率高但算力低数据复用不足块太小增大 BLOCK_M/BLOCK_N减少全局加载次数SM 占用率很低寄存器数爆了或共享内存超限降低每线程寄存器占用缩小微型 tile结果随机出错共享内存数据竞争同步缺失检查双缓冲的 barrier 位置误差稳定偏大累加顺序问题或缩放重复分 K 段累加检查缩放因子性能提升遇到天花板已经到访存墙算力提升无意义考虑 fusion把 GEMM 与其前后算子合并4.5 一个重要的调试方法论DeepGEMM 这类低层优化最忌讳“一次改很多”。我自己的习惯是每次只改动一个参数然后跑性能回归和数值对比。如果一次改动超过两个变量出了问题根本分不清是哪个改动引起的。另外调试内核时一定要保留一个用朴素三重循环实现的参考版本放到 CPU 或 GPU 上编译运行。每次修改都能和参考版本对比正确性。这个习惯帮我躲过了很多难查的随机错误。还有一个细节共享内存的未初始化区域容易被忽略新分配出的共享内存如果不显式清零某些边界条件下读出来就是脏数据导致复现不稳定的 bug。5. 进阶扩展从单个 GEMM 内核到真实系统的跨越5.1 Epilogue 融合把下一个算子拉进内核真实模型里 GEMM 后面通常紧跟着偏置、激活函数、归一化、量化缩放这些操作。如果每个算子都单独起一个内核数据就得写回显存再读出来一次来回就是几百微秒。Epilogue 融合就是把 GEMM 输出结果在寄存器里先做完一顿操作再写回显存。我在一个推理项目里把 GEMM Bias GELU Permute 融合成一个内核最终端到端延迟减少了约 15%。这个收益已经非常可观了而且实现起来并不复杂核心变化只是在一个 kernel 的尾部修改 store 之前的计算段。如果你在做推理框架优先把 GEMM 链路上的算子融合掉性价比极高。Epilogue 融合还有一个隐蔽的好处它允许你对输出做不同的数据布局转换。比如有些硬件偏好 NCHW而 GEMM 产出的是类似 NHWC 的布局融合后的内核可以直接在寄存器里处理转换省掉独立的 layout 转换内核。5.2 Split-K 与多级并行把大 K 拆开当 M 和 N 都不大、但 K 很大的时候一个 block 在 K 维上循环很多次每次循环都有同步和加载开销。Split-K 的玩法是把 K 维切成几段让不同的 block 分别算一部分局部累加最后再做一次归约。这个技术在推理场景中非常好用因为推理的 batch 往往很小M 很小导致并行度不够SM 利用不起来。Split-K 可以用并行度换单 block 计算量把 GPU 的空闲单元填满。代价是增加一次额外的 reduce 核函数或者在一个核函数里做原子累加。Stream-K 是从 Split-K 演进的一个更灵活的调度方式。它把输出块按剩余工作量动态分配到不同 block 上让所有 block 的完成时间尽量均衡。这个概念最初是针对负载不均衡问题提出的实际使用中效果显著。如果你们的推理服务 batch 大小波动很厉害这个方向值得投入。5.3 自动调优让程序自己找最优参数固定的块大小不可能在所有形状上都最优。所以大多数 DeepGEMM 类项目会带一个自动调优层比如用编译期常量模板生成一组候选配置然后分别在目标设备上跑一遍实际形状和精度挑出最快的那组参数。自动调优需要注意两点。第一调优集不能太大否则每次启动模型都要花几分钟搜索。通常把 BLOCK_M、BLOCK_N、BLOCK_K、warp 数几个变量控制在有限组合内就够了。第二调优结果要缓存并按硬件型号和输入形状做 key。不要每次进程启动都重新调。有些框架的 JIT 编译也支持这个流程启动时自动生成并缓存优化内核。DeepGEMM 名字里的“Deep”前缀某种意义上也暗示了这套“为深度网络深度定制”的思路从框架层拿到真实形状分布据此搜索出当前设备上的最优实现。5.4 实测视角什么时候能明显看到收益我简单说一下什么条件下最能感受到 DeepGEMM 类方案的收益供你判断值不值得上。第一小 M 场景M ≤ 64通用 BLAS 经常跑不满专用内核收益巨大。第二FP8 或 INT8 推理场景专用内核可以直接处理量化缩放省掉额外的 kernel。第三算子融合场景GEMM 与前后算子融合后端到端收益比单独优化 GEMM 更明显。第四多租户或动态形状场景专用内核配合自动调优可以显著降低算子延迟波动。反过来如果你的场景是训练大模型、M 很大、形状很规整、精度使用标准混合精度那么闭源 BLAS 库已经足够高效了DeepGEMM 能带来的边际提升有限。判断是否投入关键是先算算当前性能距离硬件峰值还有多远。6. 一些小建议与个人体会最后聊点实际体会。我见过不少朋友一上来就想写底层指令集版本连原型都没跑通结果调试了几天连正确性都没保证。我的建议是先快速用高级语言验证算法思路是不是可行再决定要不要花时间手写底层。原型的性能往往已经够用只有在 profiling 明确显示某个环节存在瓶颈时才值得动手改写那部分代码。还有做性能优化时不要沉迷于“数字好看”。TFLOPS 是手段不是目的。你的模型实际延迟下降了多少显存占用是否真的降了数值误差是否被业务方接受这些才是最终要回答的问题。我自己的体会是最后 10% 的优化往往需要消耗 50% 的时间除非是竞标项目或者硬件条件逼迫否则做到 80%-90% 就收手是个性价比很高的决定。另外一个小建议保持内核的“可读性”。你优化的代码很可能在三个月后还要回来改如果全是一堆魔数magic number和层层嵌套的宏定义未来你会感谢自己当初多加了几个注释。好的 GEMM 内核既是性能尖刀也是值得反复研读的数据流教科书。动手之前多花十分钟画清楚内存布局和线程映射后面能省下几小时的排查时间。