资讯动态

DeepGEMM实战:矩阵乘法优化从朴素内核到极致性能

发布时间:2026/10/9 7:40:37 来源:尧图企业网站定制
矩阵乘法在深度学习里的地位有点像是楼房里的钢筋你看不见它但每一层都得靠它撑着。无论是卷积、全连接还是Transformer里的自注意力底层都跑在广义矩阵乘法GEMM这一个运算上。DeepGEMM这个名字听起来很朴素但背后做的事情一点都不朴素——它是一个专门面向深度学习场景、把矩阵乘做到极致性能的项目。今天这篇就从一个实战工程师的角度把这类项目拆开揉碎讲清楚为什么要做、怎么做、踩过哪些坑以及哪些地方值得直接参考照搬。这篇内容适合正在做推理引擎优化、算子库开发、AI芯片工具链或者单纯想把模型跑得更快的人。如果你只是在调库阶段也不急着划走理解了GEMM的优化思路之后你在排查模型性能瓶颈时会有完全不一样的视角。1. 为什么要专门做GEMM优化1.1 深度学习运行时时间几乎都烧在GEMM上先说一个很多刚接触性能优化的人不太注意的事实在一个典型的大语言模型推理过程中GEMM类算子能占到全模型计算量的七成以上。全连接层是GEMM卷积通过重排数据变成GEMM注意力机制里的QKV投影、Attention分数计算、输出投影本质上还是GEMM。我拿一个典型的Transformer层来估算一下。假设隐藏层维度为H序列长度为T批大小为B。QKV三个投影各有一次 GEMM维度是 (BT) x H 乘 H x H。仅这三个投影的计算量就是 3 * BTHH。再算上Attention内部的两次矩阵乘以及MLP层的升维降维这一个层里十次GEMM跑不掉。在Profiler里你会看到非常直观的结果模型整体耗时曲线中GEMM火焰图占了一大截。也就是说优化GEMM就是在优化整个模型的端到端延迟这个杠杆的效率极高。这也是DeepGEMM这类项目存在的意义——把一个算子的性能挖到极限整个模型的收益立刻体现出来。1.2 通用GEMM库的通用代价有人会说现有的深度学习框架底层不早就有GEMM库了吗为什么还要自己造轮子这里有个核心矛盾通用库要兼容不同硬件、不同形状、不同数据类型、不同精度组合因此内部充满了分支判断、形状推断、多版本选择和动态调度。这些逻辑在通用性上是优势但在性能敏感路径上就是额外开销。举个例子一个通用矩阵乘库拿到一个形状 (M4096, N4096, K4096) 的任务它要先判断用哪个分块策略、是否需要转置、用哪种指令集、是否做尾数处理然后再决定实际执行的kernel。这些判断本身不会花太多时间但对于一个只有几微秒执行时间的小GEMM来说调度开销占比就可能很可观。更麻烦的是通用库为了稳妥往往不会针对某个特定形状去做极端的分块参数搜索。DeepGEMM的路线刚好相反先接受限制然后换取性能。比如锁定某一类硬件、锁定FP16/BF16和FP32累加的组合、锁定某个形状区间这时候编译器可以提前把分支全部抹掉代码生成可以做到极其纯粹的流水线。这个取舍很关键——没有免费的午餐但你可以选择不吃通用这盘菜。1.3 什么场景下真正需要动手写GEMM我不建议任何人为了炫技去重写矩阵乘。动手之前先想清楚场景值得写的场景大概有三个场景特征收益自研推理引擎需要极致低延迟端到端微秒级优化每一微秒都能转化为业务价值AI加速芯片软件栈硬件指令需要配套的高效实现不优化算子芯片算力就白费训练框架底层算子大Shape计算密度很高算力利用率直接决定训练速度节省大量GPU算力成本反过来如果只是做应用层开发、还在用框架自带的算子拼接模型那我建议你先别碰这个领域。先把思路理解清楚等真的遇到性能瓶颈时这些知识才用得上。自研GEMM是一个高投入、高回报的领域投入之前一定要确认自己的性能目标确实无法通过常规手段达成。2. 核心设计与算法拆解2.1 内存墙才是GEMM的第一敌人矩阵乘的数学定义大家都很熟悉C[i][j] Σ A[i][k] * B[k][j]。朴素实现里每计算一个输出元素需要把A的某一行和B的某一列逐个读出来相乘累加。听起来很简单但这里藏着一个巨大的性能陷阱——内存访问次数。假设我在一个MNK4096的矩阵乘里只做最朴素的双重循环。计算总量是4096的三次方约687亿次乘加这个量对现代处理器来说不算恐怖。但问题是如果每次乘法都要从主存读两个数内存带宽完全撑不住。一个GPU的内存带宽可能是每秒几百GB甚至更高但每秒能搬运的元素数量和这687亿次乘法需要的访存量相比差出几个数量级。这就是所谓的内存墙。GEMM优化的核心目标就是把数据从内存搬到寄存器后尽量多复用几次。搬运一次数据参与越多乘法运算越好。数学上这叫算术强度单位是FLOPs/Byte。算术强度越高越能逼近计算峰值而不是带宽瓶颈。2.2 三级分块从主存到寄存器的接力赛真正的高性能GEMM实现几乎都遵循同一个套路三级分块。这个思路用大白话解释就是——不一次算完整个矩阵而是把一个巨大的乘法拆成很多小方块让每一块数据在缓存和寄存器里被充分反复使用。第一级面向主存和L2缓存做分块。矩阵太大时整块读入不现实先把A按行方向、B按列方向切成较大的块例如128x256的尺度保证数据块能装进L2缓存。块在缓存里被读取时后续的乘法就能从L2而不是主存取数。第二级面向共享内存或片上SRAM做分块。在GPU这类处理器上每个线程块通常有一块高速共享内存。二级分块把L2里的数据进一步切成适合共享内存的尺寸例如64x64或64x32。由于共享内存的带宽远高于主存这一步能大幅加快数据喂给计算单元的速度。第三级面向寄存器做分块。这是真正决定计算效率的一级。每个线程要算的并不是一个点而是一个小块比如单个线程覆盖C矩阵的一个4x4或8x8小区域。这样线程从寄存器读取A的一小段可以同时与B的多个值相乘数据复用率会大幅度提升。用快递分拣来类比主存是仓库L2缓存是转运中心共享内存是小区驿站寄存器是你手里的快递车。每一次从更远的层级取数据时间成本都高一个量级所以你要尽量让数据在最近的层级完成最多的交接。2.3 B矩阵的布局变换能省掉多少事GEMM还有一个经常被忽视的细节矩阵的存储布局。大多数框架用行主序存矩阵也就是A[row][k]在内存里连续。A按行读取时访问很友好但B就麻烦了因为计算C[i][j]时需要按列读取B[k][j]。列访问在行主序存储里意味着每次跳跃一个K的步长缓存命中率很差。所以高性能实现里有一个标准操作先把B矩阵打包成连续的数据块甚至做一次隐式转置。这样计算核心读取B时内存访问是连续的硬件向量化指令也更容易发挥。深度学习的推理场景里这个操作有一个特殊优势——模型的权重是固定的。权重矩阵可以在模型加载阶段就转换好转置、分块、打包甚至量化格式转换一并完成。运行时GEMM直接消费这个预处理过的权重等于把布局转换的成本从每次推理挪到了启动阶段这对低延迟在线推理的意义非常实际。2.4 混合精度既要速度也要稳定现代AI硬件普遍支持低精度数据类型比如FP16、BF16甚至FP8。低精度数据的数值范围有限但乘法的吞吐量通常远高于FP32。于是常用的方案是计算输入用低精度存储和读取乘法在硬件低精度单元上执行累加器则使用FP32甚至更高精度。为什么累加器一定要用FP32因为矩阵乘K维往往很大动辄上千甚至几千。如果每一步累加都用低精度误差会像滚雪球一样积累。以FP16为例它的尾数只有约11位有效精度累加几千个数值之后误差完全不可控。BF16虽然指数范围很大但尾数只有7位精度更差。累加用FP32是最基本的底线。实际代码里通常长这样half a_val A[k * K i]; // 低精度输入 half b_val B[k * N j]; float acc fmaf((float)a_val, (float)b_val, acc); // FP32累加不要图省事直接用低精度累加这个坑后面在排查实录里还会专门说。3. 从朴素实现到一个能跑的DeepGEMM内核3.1 环境与工具链准备写GEMM内核前先把环境问题理清楚。我这里按主流GPU并行编程模型来写你如果用其他加速器思路可以原样平移。开发环境需要一个支持并行线程模型的语言工具链、GPU驱动、性能分析工具。其中性能分析工具特别重要光靠代码里打印耗时远远不够你得能看到访存吞吐、计算利用率这些内部指标。编译选项也要注意。开优化是必须的具体来说要让编译器放开手脚去重排指令、分解循环、使用向量化指令。如果编译器连基础优化都没开后面写的所有优化技巧都会打折扣。3.2 第一版朴素点积实现一切优化都从朴素实现开始。用一个线程计算一个输出元素最直接__global__ void gemm_naive(const float* A, const 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 acc 0.0f; for (int k 0; k K; k) { acc A[row * K k] * B[k * N col]; } C[row * N col] acc; } }这个版本能跑但性能一塌糊涂。原因前面已经埋下伏笔每个线程每做一次乘法就要从全局内存读两个数。矩阵稍微大一点内存带宽就会耗尽。我在实测一个1024三方的矩阵时这个版本只能用掉硬件峰值算力的百分之几其他时间全在等数据。它有几个缺陷值得记下来全局内存访问完全不连续B的列访问尤其致命每个A元素被反复读取没有任何复用线程之间没有协同纯粹各算各的完全没有利用向量化或张量指令3.3 第二版让每个线程算一个子块这一版的核心改动是数据复用。让每个线程负责C矩阵中一个TILE_M x TILE_N的小块而不是单个点。比如每个线程算4x4那么它读取一个A元素时可以同时服务4个输出读取一个B元素时也可以服务另外4个输出。相当于把访存量直接除以TILE_N或TILE_M。#define TILE_M 4 #define TILE_N 4 __global__ void gemm_tile(const float* A, const float* B, float* C, int M, int N, int K) { int base_row (blockIdx.y * blockDim.y threadIdx.y) * TILE_M; int base_col (blockIdx.x * blockDim.x threadIdx.x) * TILE_N; float acc[TILE_M][TILE_N] {0.0f}; for (int k 0; k K; k) { for (int i 0; i TILE_M; i) { for (int j 0; j TILE_N; j) { acc[i][j] A[(base_row i) * K k] * B[k * N base_col j]; } } } for (int i 0; i TILE_M; i) { for (int j 0; j TILE_N; j) { C[(base_row i) * N base_col j] acc[i][j]; } } }这版的性能会比朴素版提升好几倍但仍然不是最优。它的局限在于A和B还是直接从全局内存读没有经过高速存储层级。下一步要做的是把整块数据搬到共享内存里让线程块内所有线程共享。3.4 第三版共享内存协同加载这一步是多数高性能GEMM的标配线程块先把要用的A块和B块搬进共享内存计算时只从共享内存读取。共享内存的带宽比全局内存高一到两个数量级代价是容量有限且需要手动管理。我用的典型分块尺寸是 BM64、BN64、BK16。每个线程块负责C矩阵的64x64区域内层循环每次加载A的64x16和B的16x64到共享内存计算64x64的输出块。代码骨架#define BM 64 #define BN 64 #define BK 16 __global__ void gemm_shared(const float* A, const 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[4][4] {0.0f}; for (int k0 0; k0 K; k0 BK) { // 协作加载A和B的块 for (int idx threadIdx.x; idx BM * BK; idx blockDim.x) { int i idx / BK; int k idx % BK; As[i][k] A[(block_row i) * K k0 k]; } for (int idx threadIdx.x; idx BK * BN; idx blockDim.x) { int k idx / BN; int j idx % BN; Bs[k][j] B[(k0 k) * N block_col j]; } __syncthreads(); // 每个线程计算4x4子块 for (int k 0; k BK; k) { for (int i 0; i 4; i) { for (int j 0; j 4; j) { int row (threadIdx.y * 4 i); int col (threadIdx.x * 4 j); acc[i][j] As[row][k] * Bs[k][col]; } } } __syncthreads(); } // 从寄存器写回全局内存 for (int i 0; i 4; i) { for (int j 0; j 4; j) { int row block_row threadIdx.y * 4 i; int col block_col threadIdx.x * 4 j; C[row * N col] acc[i][j]; } } }这里有几个细节值得展开。首先是__syncthreads()的妙处线程块必须保证数据全部加载完才能开始计算否则会读到脏数据同样计算完成前不能被下一轮循环覆盖共享内存。同步放在两个阶段之间这是共享内存版本的命脉。其次是4x4子块的选择。线程块64x64每个线程算4x4那么一个块需要 (64/4) * (64/4) 256 个线程这个数字很合理。子块太小复用率不够子块太大寄存器压力会爆炸。寄存器属于稀缺资源一个内核能用的寄存器数量有限每个线程超过一定数量就会降低占用率反而把并行度拖垮。3.5 再加上混合精度和形状感知前面几版都是纯FP32。深度学习里真正的加速来自精度拆解。把输入换成FP16乘加单元吞吐直接翻倍。核心改动是数据类型和转换逻辑__global__ void gemm_fp16_mixed(float* C, const half* A, const half* B, int M, int N, int K) { float acc[4][4] {0.0f}; // 共享内存用half __shared__ half As[BM][BK]; __shared__ half Bs[BK][BN]; // 结构同前但读half后立即转float累加 }形状感知是什么就是让内核能根据M、N、K的实际大小选择不同的分块参数和循环策略。比如N很小而M很大的形状传统的方形分块会让大量线程处理空区域填充浪费严重。此时可以选择窄长条分块减少空闲线程。这里给一个实操经验分块参数不要靠感觉定写一个自动调优脚本遍历可能的BM、BN、BK组合以及线程块内线程排列方式把每个组合实测一遍。很多看似合理的参数实测结果完全反直觉。比如BM64、BN64不一定比BM128、BN64好因为大块意味着同步开销和共享内存压力都上升。3.6 尾数处理与边界条件实际矩阵乘法的形状很少是分块尺寸的整数倍。末尾总会有一些边角数据。处理方式有两种一是填充把矩阵扩展到分块尺寸的整数倍多余部分填零计算照常跑完再丢弃二是边界判断在内核里加 if 语句跳过越界访问。我建议在推理权重固定的场景下尽量用填充。原因很直接边界判断是动态分支GPU遇到分支发散会冻结部分线程性能损失严重。填充虽然增加了一点点冗余计算但能让所有线程保持整齐的节奏。权重是预先知道的填充一次之后每次推理都不用重复处理。4. 实测、指标与问题排查实录4.1 用哪些指标判断内核好坏判断GEMM内核做得行不行不能只看每秒多少万亿次运算。你需要几个关键指标配合起来看。指标含义判断标准算力利用率实际FLOPs除以硬件理论峰值好的实现能到70%以上满血实现接近90%内存带宽利用率实际访存带宽除以理论带宽受带宽限制的内核应尽量接近理论值寄存器溢出是否有局部变量被压到局部内存溢出出现性能基本没救占用率活跃线程占硬件容量比例太低说明并行度不足建议每个优化版本都把这几个指标记录成表格方便横向对比。不要让工程师凭感觉说这版快了要用数据证明哪个改动真正起了作用。4.2 四个高频问题的排查实录第一个问题精度崩了。这是一个非常经典的坑。我见过有人为了性能把累加器换成了FP16结果训练出来的模型权重发生变化但推理结果完全不对。排查方法其实很简单把同一个GEMM分别用FP32和混合精度跑一遍算一下C矩阵的最大绝对误差和相对误差。误差超过千分之一就要警惕超过百分之一基本不可用。修复方案就是累加器必须用FP32同时检查有没有编译器帮你把累加降精度优化掉了。第二个问题性能掉一半但找不到原因。常见的原因是访存没有对齐。GPU对128字节对齐的访问效率极高未对齐的访问会拆成多次内存事务。很多时候矩阵的列数不是对齐因子导致每一行起始地址错位。解决办法是在矩阵行之间填充padding让每一行的起始地址对齐。第三个问题小形状时性能特别差。GEMM的启动和尾部开销在小shape下非常致命。一个kernel的执行时间是 launch 开销 实际计算时间。形状太小时启动开销和同步等待占据大头。此时要考虑算子融合——把连续的GEMM、LayerNorm、激活函数合并成一个kernel减少启动次数和数据写回。第四个问题共享内存使用过度导致占用率下降。共享内存是物理资源每个线程块用得多能同时驻留的块就少。一旦块太少就无法掩盖访存延迟。遇到这类问题先看占用率再决定是缩小BN尺寸还是调整共享内存分配策略。4.3 排查思路先算理论极限再对照实测我自己排查性能问题时一定会先把理论极限写出来。假设硬件算力是A TFLOPS那么一个需要F FLOPs的矩阵乘理论最少耗时是 F / A。如果实测耗时和理论极限差了三倍以上而且算力利用率也不高那大概率是访存或并行度问题。先看Profiler里访存指令占比再看缓存命中率最后看分支发散情况。按这个顺序排查绝大多数问题都能定位。提示不要一上来就怀疑编译器。编译器确实会产生低效代码但多数情况问题出在访存模式和算法结构上。先确认算法没有问题再考虑手工调汇编。4.4 一些不值得踩的坑还有几条比较隐性的弯路我整理一下过度追求编译器魔法手动写一些看似聪明的指针重排之前先看看编译器生成的汇编。现代编译器对循环和指针别名分析比十几年前强太多了。忽略缓存块大小分块参数不仅要考虑共享内存容量还要考虑L2缓存大小。块太大L2装不下一轮所需的数据又回到主存取数。把所有线程都塞满线程过多时调度器开销可能超过收益。每个线程的活太少也不行得让每个线程有足够的寄存器级数据复用。5. 从指标到实战一个调优案例的复盘5.1 案例背景与初始状态在一个性能调优项目里我们用以上思路打磨一个推理引擎的GEMM算子。输入形状固定为M512、N4096、K4096数据格式为BF16输入、FP32累加。初始版本只做了最基本的向量化和简单分块算力利用率在45%左右。45%这个数字意味着什么理论上一块卡有10个单位的算力实际只用了4.5个。对一个固定形状的推理算子来说这个数字很难让人接受。5.2 三轮优化过程第一轮改的是共享内存分块。把BM从32改为64BN从32改为64BK从16改为32。华点在于让每个线程块内数据复用率提高了四倍。这一轮之后算力利用率升到58%。第二轮改的是B矩阵的布局。推理场景的权重可以离线打包我们把B提前做转置和重排让计算循环里对B的访问完全连续。这一轮效果很明显利用率从58%跳到72%。因为B的访存模式彻底变好了内存事务数量大幅下降。第三轮是微妙的手工调优。发现单个线程计算4x4子块时寄存器里A和B的读取存在冗余。改成8x8子块后A和B复用各多了一倍但寄存器占用略升。实测下来8x8子块更好利用率再升到81%。5.3 最终收益与总结三轮调整下来同样的硬件、同样的精度端到端推理时延降低了将近四成。这个过程很有代表性提升不是单点魔法而是逐层叠加的复利效应。6. 最后的落地建议与个人体会做这类算子优化最大的经验就是一次只改一个变量。把分块尺寸、数据布局、精度策略、线程组织这些因素拆开测试每改一个就记一次数据。否则同时改了三样东西性能变好你根本不知道其中哪样拖着后腿。另一个深刻的教训是拿真实业务形状去测不要用标准方阵。实际模型里的GEMM大多数是不规则形状M、N、K之间的比例五花八门。拿2048x2048方阵测出来的最优参数到真实形状上很可能不是最优。所以我建议所有调优都直接嵌入真实模型的profiling流程中。还有一点关于精度的经验混合精度调优时误差分析不能少。尤其当累计维度K很大时即便用FP32累加输入低精度的截断误差也可能影响某些对精度敏感的场景。对比基准输出和优化输出的最大误差只要数量级没有异常就可以放心用。如果你想从零开始实践我的建议路径是先做一次轮廓分析确认GEMM确实是瓶颈接着从朴素版本开始按上面步骤迭代每一步都记录指标。等你亲手把某个形状的GEMM从30%利用率调到了75%以上再回头去看那些商业库的实现思路会有一种豁然开朗的感觉。矩阵乘不是一个新问题但它永远值得我们拿出新态度去对待。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑