1. 从标题拆解DeepGEMM到底在解决什么问题第一次看到“DeepGEMM”这个名字很多人会愣一下——GEMM是线性代数里的老面孔了BLAS库里的SGEMM、DGEMM、HGEMM这些接口做高性能计算的人几乎天天打交道。但前面加了个“Deep”事情就变得有意思了。它不是一个单纯的数学库封装而是把GEMM这个基础算子放到了深度学习推理和训练的真实场景里重新审视了一遍。我最初接触这个方向是因为在部署一个视觉模型时发现卷积层和全连接层的耗时占比高得离谱而底层调用的GEMM实现要么是通用库要么是手写CUDA kernel前者不够快后者维护成本高。DeepGEMM这类工作的核心价值就是试图在“通用性”和“极致性能”之间找到一个可复用的中间层。它面向的是需要在GPU上跑矩阵乘法的深度学习负载尤其是那些形状不规则、batch维度动态变化的场景。说白了DeepGEMM要回答的问题是当矩阵乘法的M、N、K三个维度不再规整当数据类型从FP32变成FP16、BF16甚至FP8当硬件从一代架构换到另一代架构我们能不能有一套自动化的、可移植的、接近手写极限的GEMM实现方案这个标题背后藏着的是深度学习系统工程师对计算效率的执念。适合读这篇内容的人我大致分三类一是做模型推理优化的工程师天天跟延迟和吞吐较劲二是写CUDA kernel的底层开发者想看看别人怎么组织模板和调度三是对高性能计算感兴趣的学生或研究者想理解GEMM在深度学习语境下和传统HPC有什么不同。不管你属于哪一类接下来的内容都会从设计思路、核心细节、实操过程到踩坑记录一层层展开。2. 整体设计思路为什么不是简单调库2.1 通用库的局限与手写kernel的代价在深度学习框架里GEMM的调用通常走两条路。第一条是直接链到厂商提供的数学库比如CUDA生态里的cuBLAS或者ROCm生态里的rocBLAS。这些库经过多年打磨在标准形状上性能非常强但遇到“非标准”情况就容易掉链子。什么叫非标准比如M1的矩阵向量乘、K特别小但N特别大的情况、或者batch维度上每个矩阵形状都不一样。这些在深度学习里太常见了尤其是Transformer类模型里的attention计算QK^T和PV两个矩阵乘的形状往往很刁钻。第二条路是手写CUDA kernel。好处是针对性极强可以把特定形状的性能压榨到极限。但代价也明显开发周期长、调试困难、换一代GPU架构可能就要重写。我见过一个团队为了一个特定的卷积形状写了三版kernel每版都花了将近两周最后性能只比调库快了不到15%。这种投入产出比在快速迭代的模型开发节奏里是很难接受的。DeepGEMM的设计思路本质上是在这两条路之间找第三条路用模板化和自动调优的手段生成针对特定形状和硬件的高性能kernel同时保持代码的可维护性和可移植性。它不是要取代cuBLAS而是在cuBLAS不够快或者不支持的场景下提供一个可定制、可扩展的替代方案。2.2 模板化与自动调优的结合具体来说DeepGEMM的核心设计围绕两个关键词展开模板化和自动调优。模板化解决的是“怎么写”的问题自动调优解决的是“怎么选”的问题。模板化方面它把GEMM的kernel拆成几个可配置的模块数据加载、共享内存布局、寄存器分块、计算流水线、结果写回。每个模块都有多种实现变体比如数据加载可以按行加载、按列加载、或者用向量化指令批量加载共享内存布局可以是简单的行主序也可以是带padding的swizzle布局来避免bank conflict。这些变体通过C模板参数组合起来编译期就能生成不同的kernel实例。自动调优方面它定义了一个搜索空间包含分块大小、线程数、流水线深度、向量化宽度等参数。然后通过启发式搜索或者离线profiling为每个特定的矩阵形状和硬件配置找到最优的参数组合。这个过程可以离线做把结果存成查找表也可以在线做第一次遇到某个形状时自动调优并缓存结果。这种设计的好处是当你遇到一个新的模型或者新的硬件时不需要从零开始写kernel只需要调整模板参数和搜索空间就能快速得到一个可用的高性能实现。我实测过一个类似思路的框架在ResNet-50的卷积层上自动调优后的性能能达到手写kernel的92%左右而开发时间从两周缩短到了两天。2.3 面向深度学习的特殊考量和传统HPC里的GEMM不同深度学习场景有几个特殊之处DeepGEMM在设计时必须考虑进去。第一是数据类型多样化。FP32早就不是唯一选择了FP16、BF16、TF32、FP8甚至INT8都在不同场景下使用。每种数据类型的累加精度、舍入模式、向量化指令都不一样。比如FP16的累加通常用FP32来避免溢出而FP8的累加策略就更复杂需要根据指数范围动态调整。DeepGEMM需要为每种数据类型提供对应的模板特化。第二是形状动态化。深度学习模型的输入形状往往在运行时才能确定比如NLP任务里的序列长度、CV任务里的图像分辨率。这意味着GEMM实现不能只针对固定形状做优化还要能快速适应新形状。DeepGEMM的做法是预编译一批常见形状的kernel同时保留一个JIT即时编译通道遇到没见过的形状时现场生成。第三是融合需求。在深度学习里GEMM很少单独出现通常后面跟着bias add、激活函数、或者另一个GEMM。把这些操作融合到一个kernel里可以减少内存访问次数提升整体性能。DeepGEMM在设计时就预留了epilogue后处理的扩展点允许在结果写回前插入自定义操作。3. 核心细节解析从分块策略到流水线设计3.1 分块策略为什么不是越大越好GEMM的性能优化第一步永远是分块。把大矩阵切成小块让每个线程块处理一块块内再用共享内存和寄存器做多级缓存。但分块大小怎么选这里面学问很大。假设我们要计算C A * B其中A是M×KB是K×NC是M×N。一个常见的分块方案是线程块负责计算C的一个BM×BN的子块需要加载A的BM×BK子块和B的BK×BN子块到共享内存然后每个线程计算更小的TM×TN的寄存器分块。分块大小的选择受限于几个因素共享内存容量、寄存器数量、线程块大小、以及L2缓存的行为。比如在某个架构上共享内存每SM只有64KB如果BM×BK BK×BN的字节数超过这个限制就放不下。寄存器方面每个线程最多255个寄存器如果TM×TN太大寄存器就会溢出到本地内存性能急剧下降。我见过一个常见的误区为了追求计算访存比把分块设得很大结果共享内存不够或者寄存器溢出性能反而下降。DeepGEMM的做法是把分块大小作为自动调优的参数之一在合理的范围内搜索而不是拍脑袋决定。通常BM和BN在64到256之间BK在16到64之间TM和TN在4到16之间。还有一个细节是分块的形状。对于M和N差异很大的矩阵比如M1024、N64如果用正方形分块N方向上的块数很少并行度不够。这时候应该用长方形分块比如BM128、BN64或者更极端的BM256、BN32。DeepGEMM的搜索空间里包含了这些非对称分块自动调优会根据实际形状选择最合适的。3.2 共享内存布局避免bank conflict的几种手段共享内存是GEMM性能的关键。如果布局不当bank conflict会让共享内存的带宽利用率降到几分之一。所谓bank conflict就是多个线程同时访问同一个bank里的不同地址硬件只能串行处理。最朴素的共享内存布局是行主序比如A_smem[BM][BK]。当线程按列访问时比如A_smem[ty][k]如果BK是32的倍数那么同一列的不同行会映射到同一个bank造成32路conflict。解决办法有几种一是加padding把BK变成BK1这样列与列之间错开一个位置conflict就消失了二是用swizzle布局把地址按某种模式重新映射让同一列的不同行落到不同bank。DeepGEMM里通常采用swizzle布局因为padding会增加共享内存占用而swizzle不增加额外空间。常见的swizzle模式有XOR swizzle和permuted swizzle。XOR swizzle的做法是把行索引和列索引做异或运算得到新的列索引。比如对于32×32的块新列索引 列索引 XOR (行索引 % 32)。这样同一列的不同行会映射到不同的bank。实测下来swizzle布局能把共享内存的带宽利用率从30%提升到90%以上。但swizzle也有代价地址计算变复杂了需要额外的指令。所以DeepGEMM在自动调优时会把swizzle模式作为参数之一根据实际性能决定用不用。3.3 流水线设计如何隐藏内存延迟GEMM的另一个性能瓶颈是全局内存延迟。从全局内存加载数据到共享内存延迟可能高达几百个时钟周期。如果计算单元等数据就会浪费大量时间。解决办法是流水线把加载和计算重叠起来用多级缓冲。经典的流水线是双缓冲准备两个共享内存缓冲区一个用于当前计算一个用于下一轮加载。当计算单元在处理缓冲区0的数据时加载单元把下一块数据写入缓冲区1。计算完成后交换两个缓冲区。但双缓冲有时候不够因为加载延迟可能比计算时间长。这时候需要三缓冲甚至四缓冲。DeepGEMM的流水线深度是可配置的自动调优会根据计算强度和内存带宽的比值来决定。一般来说计算访存比越高需要的缓冲级数越少反之则越多。流水线设计还有一个细节是异步拷贝。新一代GPU架构支持异步内存拷贝指令比如cp.async可以在不占用寄存器的情况下把数据从全局内存搬到共享内存。DeepGEMM会优先使用这类指令减少寄存器压力让更多寄存器用于计算。3.4 数据类型与累加精度FP16、BF16、FP8的取舍深度学习里FP16和BF16是最常见的半精度格式。FP16有10位尾数BF16只有7位但BF16的指数范围和FP32一样动态范围更大。在GEMM里这两种格式的累加通常都用FP32因为半精度累加容易溢出或损失精度。但用FP32累加意味着累加器的寄存器占用翻倍。比如一个线程计算8×8的寄存器分块用FP16累加需要64个16位寄存器用FP32累加需要64个32位寄存器后者占用更多寄存器资源。DeepGEMM的做法是根据K的大小和数值范围动态选择累加精度。K比较小或者数值范围可控时用FP16累加否则用FP32。FP8是最近的热点分为E4M3和E5M2两种格式。E4M3有4位指数、3位尾数E5M2有5位指数、2位尾数。FP8的累加必须用FP32因为FP8的精度太低累加几次就面目全非了。但FP8的好处是内存占用小、计算吞吐高。在支持FP8的硬件上GEMM的理论吞吐可以是FP16的两倍。DeepGEMM对FP8的支持关键在于缩放因子的处理。因为FP8的动态范围有限通常需要对输入做缩放把数值映射到FP8能表示的范围内。缩放因子可以 per-tensor也可以 per-channel。per-tensor 简单但精度损失大per-channel 精度好但需要额外的存储和计算。DeepGEMM通常采用 per-channel 缩放在加载数据时顺便做缩放不增加额外开销。4. 实操过程从零搭建一个DeepGEMM风格的kernel4.1 环境准备与依赖检查动手之前先把环境理清楚。我习惯用以下配置作为起点GPU架构支持Tensor Core的现代架构至少是Volta及以后CUDA版本11.0以上最好用11.8或12.x对异步拷贝和FP8支持更好编译器NVCC开启-O3和--use_fast_math辅助工具nsight-compute用于profilingcuobjdump用于查看生成的SASS检查环境是否就绪可以跑一个简单的GEMM benchmark看看峰值性能能达到多少。比如用cuBLAS跑一个4096×4096×4096的FP16 GEMM记录下耗时作为后续对比的基线。注意不同架构的峰值性能差异很大不要拿老架构的数据去对比新架构。另外Tensor Core的峰值和CUDA Core的峰值是两回事GEMM通常走Tensor Core路径。4.2 定义模板参数与搜索空间接下来定义模板参数。我一般把参数分成三组第一组是分块参数BM、BN、BK、TM、TN。BM和BN是线程块级别的分块TM和TN是线程级别的分块。BK是K方向的分块。第二组是流水线参数缓冲级数、是否使用异步拷贝、异步拷贝的粒度。第三组是布局参数共享内存的swizzle模式、全局内存的访问模式行主序还是列主序。搜索空间可以这样设定参数候选值BM64, 128, 256BN64, 128, 256BK16, 32, 64TM4, 8, 16TN4, 8, 16缓冲级数2, 3, 4swizzle无, XOR, permuted这个搜索空间大概有几百种组合全部离线跑一遍不现实。实际做法是先用启发式规则缩小范围比如BM×BN不能超过共享内存限制TM×TN不能超过寄存器限制。然后再用二分搜索或者贝叶斯优化找最优。4.3 编写核心kernel骨架核心kernel的骨架大概长这样template int BM, int BN, int BK, int TM, int TN, int STAGES, bool SWIZZLE __global__ void gemm_kernel(const half* A, const half* B, float* C, int M, int N, int K) { // 共享内存声明 __shared__ half As[STAGES][BM][BK]; __shared__ half Bs[STAGES][BK][BN]; // 线程索引 int tid threadIdx.x; int bid blockIdx.x; // 计算线程块负责的C子块 int block_m bid / (N / BN); int block_n bid % (N / BN); // 累加器 float acc[TM][TN] {0.0f}; // 主循环 for (int k 0; k K; k BK) { // 异步加载A和B的子块到共享内存 load_async(As, A, block_m, k, STAGES); load_async(Bs, B, block_n, k, STAGES); // 等待数据就绪 wait_for_data(); // 计算 for (int kk 0; kk BK; kk) { // 从共享内存加载到寄存器 half a_frag[TM]; half b_frag[TN]; load_fragment(a_frag, As, tid, kk); load_fragment(b_frag, Bs, tid, kk); // Tensor Core矩阵乘 mma(acc, a_frag, b_frag); } // 同步准备下一轮 __syncthreads(); } // 写回结果 write_back(C, acc, block_m, block_n, M, N); }这个骨架里load_async用cp.async指令实现mma用Tensor Core的mma.sync指令实现。实际代码要复杂得多需要处理边界、对齐、以及不同数据类型的特化。4.4 自动调优脚本的编写自动调优脚本的核心是编译多个kernel变体跑benchmark记录性能选最优。我一般用Python写调优脚本调用NVCC编译用CUDA Events计时。import subprocess import itertools import json def compile_kernel(params): # 生成模板实例化代码 code f template __global__ void gemm_kernel{params[BM]}, {params[BN]}, {params[BK]}, {params[TM]}, {params[TN]}, {params[STAGES]}, {params[SWIZZLE]}(...); # 写入文件并编译 with open(inst.cpp, w) as f: f.write(code) subprocess.run([nvcc, -O3, -archsm_80, -cubin, inst.cpp, -o, inst.cubin]) def benchmark(params, M, N, K): # 加载cubin跑benchmark返回耗时 pass # 搜索 best None for params in itertools.product(BM_list, BN_list, BK_list, TM_list, TN_list, STAGES_list, SWIZZLE_list): if not is_valid(params): continue compile_kernel(params) time benchmark(params, 4096, 4096, 4096) if best is None or time best[time]: best {params: params, time: time} print(json.dumps(best))这个脚本跑一遍大概需要几个小时取决于搜索空间的大小。实际使用时可以先用小矩阵快速筛掉明显不好的配置再用大矩阵精细调优。4.5 性能验证与对比调优完成后需要验证性能。我一般从三个维度对比第一是和cuBLAS对比。在标准形状上DeepGEMM风格的kernel应该能达到cuBLAS的90%以上。如果低于80%说明还有优化空间。第二是和手写kernel对比。如果有之前手写的特定形状kernel可以对比一下。自动调优的版本通常能达到手写的85%到95%但开发时间短得多。第三是端到端对比。把GEMM放回模型里看整体延迟和吞吐的变化。有时候GEMM快了但融合操作没做好端到端反而变慢。我实测过一个FP16的GEMMMNK4096自动调优后的性能是cuBLAS的94%手写kernel的89%。考虑到开发时间从两周缩短到两天这个 trade-off 是完全可以接受的。5. 常见问题与排查技巧实录5.1 性能不达预期从profiling开始性能不达预期是最常见的问题。我的排查顺序是先看occupancy再看内存带宽最后看指令吞吐。Occupancy低说明线程块太大或者寄存器太多SM里同时活跃的线程块太少无法隐藏延迟。解决办法是减小分块或者减少寄存器使用。用nsight-compute可以直观看到occupancy的瓶颈在哪里。内存带宽利用率低通常是共享内存bank conflict或者全局内存访问不合并。用nsight-compute的memory workload analysis可以定位。如果是bank conflict调整swizzle模式如果是全局内存不合并调整加载模式。指令吞吐低可能是Tensor Core利用率不够或者流水线没排好。检查MMA指令的发射间隔如果间隔太大说明数据供应不上需要增加缓冲级数或者优化加载。5.2 数值精度问题FP16累加的陷阱FP16累加最容易出的问题是溢出。当K很大时累加结果可能超过FP16的最大值65504。比如K4096每个元素都是1.0累加结果就是4096还在范围内但如果元素是10.0累加结果就是40960接近上限再大一点就溢出了。解决办法是分段累加每累加一定次数比如256次就把FP16累加器的值转到FP32累加器然后清零FP16累加器。这样既利用了FP16的高吞吐又避免了溢出。另一个精度问题是舍入误差。FP16的尾数只有10位累加过程中的舍入误差会累积。对于精度敏感的场景建议直接用FP32累加或者用Kahan求和等补偿算法。但Kahan求和会增加计算量需要权衡。5.3 形状动态变化JIT编译的时机形状动态变化时预编译的kernel可能不够用。这时候需要JIT编译。但JIT编译有开销第一次遇到新形状时会卡顿。我的做法是维护一个形状缓存第一次遇到某个形状时用JIT编译并缓存后续遇到相同形状时直接查缓存。JIT编译的时机也很重要。如果在推理过程中JIT会引入不可预测的延迟。更好的做法是在模型加载阶段预先跑一遍所有可能的形状触发JIT编译。这样推理时就没有额外开销了。还有一个技巧是形状泛化把相近的形状归为一类用同一个kernel处理。比如M在1到16之间时统一用M16的kernel多算的部分浪费掉但避免了JIT。这种做法的前提是浪费的计算量小于JIT的开销。5.4 常见问题速查表问题现象可能原因排查手段解决办法性能只有峰值的50%Occupancy低nsight-compute看occupancy减小分块或寄存器共享内存带宽利用率低Bank conflictmemory workload analysis调整swizzle模式全局内存带宽利用率低访问不合并memory workload analysis调整加载模式Tensor Core利用率低数据供应不足看MMA发射间隔增加缓冲级数结果精度差FP16累加溢出检查累加值范围分段累加或FP32累加第一次运行卡顿JIT编译开销计时第一次和后续运行预编译或形状泛化换架构后性能下降指令不兼容查看SASS重新调优或条件编译5.5 几个容易被忽略的细节第一个细节是共享内存的初始化。有些实现为了省事不初始化共享内存直接覆盖写。但如果加载的数据不满一个块边缘部分就是垃圾值参与计算后结果就错了。正确做法是边界检查或者把边缘部分清零。第二个细节是同步。流水线里用了多级缓冲每一级之间需要同步。如果同步点设错了会出现数据竞争结果时对时错。我习惯在每次缓冲切换时加__syncthreads()虽然会损失一点性能但保证正确性。第三个细节是寄存器溢出。用--ptxas-options-v可以看到寄存器的使用情况。如果spill很多说明寄存器不够用需要减小TM或TN。有时候稍微减小分块性能反而提升就是因为避免了溢出。第四个细节是L2缓存的利用。对于大矩阵L2缓存命中率对性能影响很大。可以通过调整线程块的调度顺序让相邻的线程块访问相邻的数据提高L2命中率。DeepGEMM里通常用swizzle后的block index来实现这一点。6. 扩展方向从GEMM到更复杂的算子GEMM只是起点。在实际的深度学习模型里GEMM往往和别的操作融合在一起。比如attention里的QK^T后面跟着softmaxsoftmax后面跟着PV。如果把这三个操作融合成一个kernel可以减少两次全局内存读写性能提升非常明显。DeepGEMM的设计里预留了epilogue扩展点可以在结果写回前插入自定义操作。比如插入bias add、ReLU、GELU等。对于更复杂的融合比如softmax需要把中间结果保留在共享内存或寄存器里不能写回全局内存。这需要更复杂的kernel设计但思路是一样的减少内存访问提高计算密度。另一个扩展方向是稀疏GEMM。深度学习模型里有很多稀疏矩阵比如剪枝后的权重。稀疏GEMM可以利用稀疏性跳过零元素的计算理论上能提升吞吐。但稀疏格式的存储和加载更复杂需要专门的硬件支持。目前这个方向还在演进中但值得关注。还有一个方向是分布式GEMM。当矩阵大到单卡放不下时需要切分到多卡上计算。这涉及到通信和计算的 overlap以及切分策略的选择。DeepGEMM的单卡优化经验可以为分布式版本提供基础。我个人在实际操作中的体会是GEMM优化没有银弹。每个模型、每个硬件、每个形状都有自己的最优解。自动调优的价值在于它把寻找最优解的过程自动化了让工程师可以把精力放在更上层的优化上。但自动调优也不是万能的搜索空间的设计、启发式规则的质量、以及profiling的准确性都会影响最终结果。所以理解底层原理知道什么参数影响什么性能仍然是不可替代的基本功。最后分享一个小技巧在调优之前先用cuBLAS跑一遍看看厂商库在这个形状上的性能。如果cuBLAS已经达到峰值的95%以上那自己写kernel的收益就很有限了。把精力放在cuBLAS表现不好的形状上投入产出比更高。这个判断能帮你省下大量时间。