ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

DeepGEMM:深度学习矩阵乘法算子优化实战

DeepGEMM:深度学习矩阵乘法算子优化实战 我一直觉得做AI训练的如果没把一个矩阵乘法算子从底往上啃一遍就很难真正理解“深度学习为什么吃GPU算力”这句话的分量。DeepGEMM这个名字拆开来看就是深度学习和通用矩阵乘法的结合把深度学习里最核心的那个GEMM算子在硬件层面做到又快又稳。这篇文章就围绕这个主题展开讲讲GEMM为什么值钱、一个高性能GEMM算子到底要解决什么问题以及我在搭这类算子时实际踩过的坑和摸索出来的优化思路。想搞懂AI框架底层的算法工程师或者准备优化算子、冲性能的开发者这篇应该对你有用。1. GEMM为什么值钱先搞清楚矩阵乘法的江湖地位1.1 所谓“通用矩阵乘法”到底在算什么GEMM四个字母来源于General Matrix Multiplication本质就是C A × B这条最基础的计算规则。A的维度是M×KB的维度是K×NC就是M×N中间那个K是求和维度。算一遍要做M×N×K次乘法和同样数量的加法总共约2×M×N×K次浮点操作。拿一个不算夸张的例子来说一个尺寸为4096×4096×4096的矩阵乘大约要算1370亿次浮点操作。这个数字有多大呢普通桌面级处理器每秒能跑的浮点指令量连这个算例的零头都摸不到也只有现代加速卡上的大量并行计算单元能在毫秒级把它吃掉。这里有一个关键点容易被忽视GEMM不只是“会算就行”它得“被算得快”。因为深度学习里的GEMM往往不是单独出现一次而是作为一个基础算子被反复调用。一旦GEMM慢了整个模型的训练和推理都会被拖住。这也是为什么很多计算库底层的性能优化核心都压在这一个算子上。1.2 深度学习里的GEMM远比你想的更普遍很多人以为深度学习里面只有全连接层才叫矩阵乘法实际上卷积、注意力、甚至一些图神经网络里的消息传递最后都能转成GEMM。全连接层不用说了输入矩阵乘上权重矩阵就是GEMM。卷积层则有两条常见的路线一种是把输入图像按滑动窗口展开成矩阵再和卷积核矩阵做乘法叫im2col另一种更聪明的做法是隐式GEMM通过特殊的索引把卷积转化成矩阵乘不显式展开既省内存又容易复用高速缓存。注意力机制里更明显Query和Key的转置相乘得到注意力分数这里就是一个GEMM注意力分数再和Value相乘又是一个GEMM。Transformer结构里绝大部分计算量都是由这两类矩阵乘贡献的。所以你会发现一个深度学习模型的性能上限很大程度上取决于其所依赖的GEMM算子性能上限。算子本身优化到位整个模型自然受益算子写得很粗糙上面框架再优化也救不回来。1.3 为什么GEMM性能几乎决定了AI框架的成败做一个简单的比喻把加速卡想象成一个顶级厨房算力单元是厨房里的灶台芯片和显存之间的带宽是传菜通道。你菜刀再快、灶台再多如果食材输送不及时整个厨房还是只能干瞪眼。GEMM优化解决的核心问题就是怎么配合好灶台和传菜通道。业界评估训练效率时常用到MFUModel FLOPs Utilization这个指标意思是模型在加速卡上实际跑出来的算力除以加速卡标称的峰值算力。MFU越接近1说明设备利用得越充分。如果一个模型里面充满各种低效的GEMMMFU就很难看训练时间长、成本高这是所有做大规模训练的团队都不想看到的。所以“DeepGEMM”这个标题背后真正的野心并不只是简单实现一个矩阵乘法而是想方设法把这个算子挤到硬件极限。这个极限不仅取决于算力有多高还取决于有没有把内存访问、指令调度、精度转换这些细节全部拿捏住。2. 拆解DeepGEMM的核心设计思路让数据流动起来2.1 核心矛盾算力远远跑不过访存设计一个高性能GEMM算子第一步不是去看怎么写并行代码而是先看手头的硬件账本。现代加速卡都有个明显特点计算速度的增长远远快过显存带宽的增长。可以查一组宏观数据当前主流的专业加速卡FP16稠密算力往往能达到数百甚至上千TFLOPS而显存带宽还停留在几TB每秒的级别差距接近两个数量级。这里引入一个关键概念叫计算强度Arithmetic Intensity定义是“每从内存取一个字节能支撑多少次计算”。要喂饱算力计算强度必须超过一个阈值阈值约等于峰值算力除以带宽。假设某张卡的FP16算力是900 TFLOPS带宽是3TB/s那么阈值就是900×10^12 / (3×10^12) 300。意思是每从显存读入一个字节必须至少支撑300次浮点运算才算把这台机器的计算潜力吃干净。这个数字放在朴素版本的三重循环里是达不到的。因为每次计算C[i][j]都要去显存里重新拉A的第i行和B的第j列等于是每做一次乘法都重复搬运数据算术强度极低。优化的方向只有一个想尽办法增加数据复用让同一个数被多次计算利用。2.2 分块Tiling与循环重排把数据留在离计算最近的地方存储器层级像一座金字塔寄存器在塔尖容量最小但速度最快再往下是高速缓存、共享内存再往下才是显存。速度快的存储容量小速度慢的存储容量大所以高性能计算的核心原则就是让数据尽可能地待在塔尖附近的存储里而不是反复去最底层取。GEMM上的经典做法是分块。把A矩阵按M维度切成BM高的小块把B矩阵按N维度切成BN宽的小块K维度每次取BK。每一个计算块只负责一小块输出例如C的一个BM×BN子块由当前线程块独立完成。这样的话A块可以被反复用来和B的每一小块做运算数据从全局内存到共享内存只需要搬运一次之后一直在共享内存和寄存器里打转复用率就大大提升了。如果挑一个常见的分块尺寸比如BM128BN128BK16那么每个线程块从显存取2048个A元素和2048个B元素却能计算出128×128共16384个输出元素。相对朴素版本显存访问量降低到原来的几十分之一。这也是很多高性能GEMM库的标准思路看似不起眼却是性能起飞的第一步。2.3 寄存器级切分与张量核心的配合分块优化只能解决从显存到共享内存这一层真正想要逼近峰值算力还要注意到现代加速卡都有专门的矩阵乘累加单元我们习惯叫它张量核心。这类硬件单元天生就是为小规模矩阵乘准备的比如在一个指令周期里完成某个固定形状的A块和B块相乘并累加到C块上效率远超一堆通用的乘加指令。要用好张量核心线程块内还要继续往下拆把BM×BN的输出块切给多个线程束每个线程束再切给多个线程由每个线程在自己的寄存器里维护若干份输出累加器。经典的做法是让线程束里每个线程负责一个小的矩阵片段通过硬件指令一次性完成一小块乘累加。这种切分和调度策略决定了张量核心是满负荷运转还是有一半时间在空转。这里分享一个很容易被新手忽略的点如果只把A和B放进共享内存却让线程自己拿普通指令去四重循环计算张量核心完全派不上用场性能大概率只有峰值算力的零头。所以DeepGEMM这类算子的关键不在于“有没有分块”而在于“有没有把寄存器级的数据布局和张量核心的形状要求对齐”。3. 精度策略DeepGEMM为什么敢在低精度上做文章3.1 FP16/FP32混合精度是底线讨论GEMM性能时精度选择是不能绕开的。因为很多硬件对低精度运算提供的算力远远高于高精度而且低精度数据搬移量更小一个字节可以装更多数据带宽压力也更小。FP16就是深度学习里用得最多的一种低精度格式。FP16和FP32最大的区别在于位宽减半但指数位和尾数位的分配都会变化。这会导致两个问题一是尾数位变少表示精度更粗二是指数范围变窄数值很容易溢出或者下溢。所以直接拿FP16替换FP32从训练角度容易出现损失函数震荡或收敛变慢的情况。实践中我的做法是混合精度模型的权重和梯度保留一份FP32副本真正进入GEMM计算的输入数据用FP16在GEMM内部累加时用FP32精度计算完再存回FP16。这样既保住训练稳定性又享受了低精度带来的吞吐优势。这个思路也是很多训练框架的默认做法属于底线方案。3.2 FP8更低精度、更高吞吐但必须有补偿机制随着模型规模越涨越猛FP16也开始显得奢侈。新一代硬件往FP8方向靠近FP8的位宽只有8比特数据量进一步减半。FP8又分两种格式E4M3有4位指数和3位尾数动态范围相对小但精度略高E5M2有5位指数和2位尾数动态范围大但尾数精度低。通常矩阵乘法主计算会优先选E4M3梯度传播或需要大动态范围的场景用E5M2。FP8最大的红利是在同一张卡上一个计算周期能完成比FP16更多的小规模矩阵乘吞吐量几乎翻倍。但位宽减少也会让误差急剧放大特别是数据里出现一些离群的大数保存到FP8里很容易直接把整体数值范围撑爆。所以FP8 GEMM的基本配套是缩放。最常用的一种做法是块级缩放把输入矩阵切成若干个小块统计每个块内的绝对最大值然后统一除以这个最大值再转成FP8使大部分数值能落在FP8最舒适的表示区间。这个缩放因子要保持到计算结果累加之后再用回来。块切得越小缩放越适配局部数值分布但缩放因子的存储和计算开销也会越高这里有性价比的权衡。3.3 量化误差控制的实际经验我见过不少项目在FP8上翻车典型表现是Loss曲线看起来没问题但某些算子输出的梯度和FP32参考值差异非常大最终导致模型某个层彻底学不进去。问题往往出在“只看Loss不看逐层数值分布”上。可靠的做法是在切换低精度GEMM时至少要把输入、输出、中间累加结果都dump出来和FP32版本做逐元素对比。看两个指标一个是最大绝对误差一个是相对误差。如果相对误差在1e-3量级以下通常可以接受超过1e-2就要小心因为误差在经过深层网络传播后会被放大。还有一个小技巧是“延迟缩放”。每次计算都重新统计块内最大值熟悉矩阵分布后再使用上次的比例因子省掉很多额外的全局归约操作既能让数值稳定又能省时间。这类细节在普通文档里很难找到都是实打实调出来的经验。4. 从理论到实现一个可复现的DeepGEMM优化流程4.1 先写一个正确的朴素版本前面讲了那么多理论落到代码上还是要一步一步来。我建议无论你目标性能多高都要先写一个正确、朴素的GEMM版本主要目的是用来对比验证后续优化版本的正确性。朴素版本的逻辑就像老师上课教的矩阵乘法三重循环每个线程负责输出矩阵中的若干个元素// 朴素GPU版GEMM示意每个线程计算C中一个元素 __global__ void gemm_naive(float* A, float* B, float* C, int M, int N, int K) { int row blockIdx.y * blockDim.y threadIdx.y; int col blockIdx.x * blockDim.x threadIdx.x; if (row M col N) { float sum 0.0f; for (int k 0; k K; k) { sum A[row * K k] * B[k * N col]; } C[row * N col] sum; } }这个版本通常只能跑出峰值算力的几个百分点因为每个输出元素都要重复从显存拉取大量数据带宽早早就成了瓶颈。但它最大的价值是提供一个“标准答案”。之后每写一版优化我都会拿它做allclose校验如果误差超阈值说明优化实现里引入的数据布局或累加顺序出了问题。4.2 引入共享内存分块后的性能跃升朴素版本验证正确后第二步就是引入分块把A和B的块搬进共享内存。这里我给出一个常见的设计每个线程块负责BM×BN的输出块循环沿着K维度每次处理BK。// 分块GEMM示意利用共享内存提高数据复用 #define BM 128 #define BN 128 #define BK 16 __global__ void gemm_tiled(float* A, float* B, float* C, int M, int N, int K) { __shared__ float As[BM][BK]; __shared__ float Bs[BK][BN]; int block_row blockIdx.y * BM; int block_col blockIdx.x * BN; float acc[BM / 8][BN / 8] {0.0f}; // 线程内累加器示意 for (int k0 0; k0 K; k0 BK) { // 协同加载A和B的分块到共享内存 for (int i threadIdx.x; i BM * BK; i blockDim.x) { int row i / BK; int col i % BK; As[row][col] A[(block_row row) * K (k0 col)]; } for (int i threadIdx.x; i BK * BN; i blockDim.x) { int row i / BN; int col i % BN; Bs[row][col] B[(k0 row) * N (block_col col)]; } __syncthreads(); // 线程块内的乘累加逻辑这里省略张量核心映射 for (int i 0; i BM; i) { for (int j 0; j BN; j) { // 累加As和Bs } } __syncthreads(); } // 写回C }分块之后显存访问量已经大幅下降但离峰值算力还有一段距离。此时要开始考虑双缓冲也就是用两块共享内存交替装载下一轮数据当前一轮的计算和下一轮的加载同步进行。这样能隐藏全局内存访问的等待时间让计算单元在加载数据时依然有事可做。4.3 记忆屏障与流水线重叠的关键细节引入共享内存后最容易被坑的一步是同步。共享内存是同一线程块内所有线程共享的好处是方便协作坏处是线程之间访问共享内存时必须保证数据已经写好了。__syncthreads()这个栅栏操作会强制所有线程等待如果数量使用过多整个流水线会被反复拖慢。现代加速卡提供了一种异步拷贝机制可以绕过寄存器直接从显存把数据搬进共享内存并在等待拷贝完成的同时继续执行当前计算。这就是硬件级的流水线重叠。具体做法是在循环迭代k的时候预取第k1轮需要的A和B分块再同步当前计算。我发现很多新手第一次写双缓冲时把预取的代码和同步的代码顺序搞错导致每轮都死等性能没有任何提升。这类优化一定要借助性能分析软件来指导。可以看到指令等待周期占比、共享内存访问冲突次数、全局内存吞吐量这些指标数据说话比经验瞎猜可靠得多。5. 实测过程中的瓶颈分析与排查实录5.1 常见瓶颈bank conflict、尾效应、寄存器溢出优化到中后期你会发现每一个小细节都在影响最终性能其中有三个问题出现频率最高。第一是共享内存的bank conflict。共享内存物理上被划分成若干bank如果多个线程同时访问同一个bank的不同地址访问就会被串行化处理属于“抢同一个仓库的通道”。举例来说如果数组每行长度设计成128而线程束的访问模式刚好让它们每一拍都打到同一个bank上那共享内存访问速度会直接崩塌。解决办法通常是给数组加pad把每行长度改成129错开地址映射往往见效极快。第二是尾效应。输入矩阵的维度不一定刚好是分块大小的整数倍。比如M3333你用128做分块最后肯定会剩下一个不满128的尾巴。要么填充矩阵改成对齐要么在核函数里加边界判断。边界判断会增加指令开销但能避免访问越界。实测中我倾向于padding因为如果填充不影响计算结果它对计算流程的干扰最小。第三是寄存器溢出。为了保持张量核心忙碌每个线程会尽量多分配一些寄存器做累加器。但寄存器的总量有限分配太多会把中间变量挤到本地内存里本地内存实际上还是显存速度掉到十分之一以下。性能往往不是线性下降而是直接悬崖式跌停。解决办法是减少每个线程承担的累加器数量把线程块内线程数调多或者显式限制寄存器数量。5.2 踩坑记录与排查思路我自己在这些优化步骤上翻过不少车举几个印象深刻的例子。第一次做分块时性能不升反降。我用profiler一看共享内存bank conflict率高达90%罪魁祸首就是数组行宽选成了128的老实数值没有加padding。后来改成129后性能立刻涨了40%。这个教训让我养成了习惯每次改共享内存布局先看冲突率再决定优化方向。第二次是双缓冲写出来的结果全是错的。排查了很久发现是异步拷贝还没完成就开始读共享内存数据预取逻辑和计算逻辑之间的依赖没有同步好。这种错误在单个线程或小规模调试时很难发现只有到大矩阵计算才会露出明显的随机错误。之后我在关键位置加上等待异步拷贝完成的指令问题才消失。这类“丢失同步”的bug是写算子时最容易犯的隐蔽错误。第三次是FP8的误差排查。当时某个模型切到FP8后Loss曲线在训练中期突然出现一次尖峰然后又恢复。我一开始以为是学习率调整的问题后来逐个算子对比发现是某一层输入的数值动态范围很大而我的块级缩放因子在循环里没及时更新导致一批数据的最大值统计失真。修好缩放因子的更新策略后尖峰再也没有出现。5.3 性能数据怎么读别被加速比骗了优化完以后大家最喜欢晒的就是“比之前快了多少倍”。但我觉得做性能评估时要格外冷静不然很容易被加速比误导。第一要固定对比条件。矩阵规模、数据精度、硬件环境、缓存热不热都会影响结果。同样的算子尺寸大和尺寸小时的性能表现可能完全不同不能只挑最好看的数据说事。第二要关注绝对性能值。判断一个GEMM算子写得好不好最重要是看它在这张卡上达到了峰值的百分之多少。如果峰值是900 TFLOPS你跑了600那说明已经相当理想如果只跑了100哪怕比朴素版本快50倍离优秀也还有很大距离。第三要理解“大矩阵和小矩阵不是一回事”。小矩阵算例很容易被启动开销和调度开销主导大矩阵才能把计算流水填满。所以评估时要覆盖一组有代表性的尺寸不要只测一个正方形的大矩阵。我整理了一个示意性的性能对比表矩阵规模MNK朴素GEMM共享内存分块张量核心双缓冲10248 TFLOPS28 TFLOPS165 TFLOPS20489 TFLOPS35 TFLOPS210 TFLOPS40967 TFLOPS36 TFLOPS245 TFLOPS这类数字只是为了让初学者理解“不同优化阶段的巨大差异”具体值跟你手头的硬件和编译器版本直接相关别拿去当理论参考。真相只有一个一切以你自己机器上复现的结果为准。6. 给后来者的一些实操心得最后分享几条做这类算子优化积累下来的经验。如果你想复刻一条DeepGEMM式的优化路线我认为顺序很重要。第一步先选定一个目标shape不要一上来就想做出一个适用于所有尺寸的万能算子。绝大多数优化方案都和shape有关系先把一个常见尺寸做透再想怎么推广。第二步把正确性测试框架搭在前面。写一个朴素的CPU版本或者用现有框架的GEMM结果做基准每优化一步就跑一遍allclose。很多优化改动是细微的比如交换了累加顺序没有测试框架很容易漏掉误差问题。第三步性能分析工具要尽早用。不要等全部写完了再分析每个阶段都看一眼关键指标。我常用的检查顺序是先看全局内存吞吐再看共享内存冲突率最后看计算单元利用率。哪一项异常就沿着那个方向去定位。第四步不要害怕精度和性能一起调。有时为了压榨性能你可能会考虑把FP16换成FP8或者把分块尺寸调大。每一步都要同时观察误差变化和性能变化不能只盯速度。慢一点但是数值正确远比快很多但结果错乱更有价值。还有一个小技巧值得记下调试低精度GEMM时可以用“把输入数据乘上一个很小的缩放因子”这种办法来看误差是来自量化本身还是来自实现里的bug。如果缩小输入后误差的比例依然不变多半是算法实现里数值处理出了问题如果误差等比缩小说明量化机制本身在正常工作只是动态范围没有选好。算子优化这件事做到后期拼的不是灵光一闪而是对每一个细节的耐心打磨。一次一次地看指标、找瓶颈、修正实现把数据搬移量压到极限把同步开销减到最小让计算单元始终有活干这条路走下去你也会慢慢摸到硬件的脾气。希望这篇关于DeepGEMM的拆解思路能让你在自己的优化路上少走几步弯路。
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进