ARTICLE DETAIL

资讯详情

深耕商务建站与企业官网运营的一线实战洞察。

DeepGEMM:深度学习高性能矩阵乘法内核优化实践

DeepGEMM:深度学习高性能矩阵乘法内核优化实践 GEMM通用矩阵乘法是深度学习训练和推理里最底层的重型计算单元从全连接层到注意力分数计算本质都是在搬矩阵。很多人以为直接调官方闭源数学库就万事大吉但真到了性能优化、算子融合、低精度推理这些阶段手边能自己掌控的矩阵乘法内核反而不够用。DeepGEMM 就是我做的一个面向深度学习场景的高性能矩阵乘法内核集合目的很直接在常见形状上把算力吃透同时把 epilogue 融合、量化支持这些推理引擎真正需要的功能做进去。这篇内容适合正在做推理引擎、写自定义算子、或者想搞懂矩阵乘法性能瓶颈的朋友我尽量把从分块策略到 Tensor Core 指令、再到精度对齐的整个思路讲清楚。1. 为什么深度学习场景需要一份专属GEMM而不是直接调官方库1.1 GEMM 是深度学习的发动机再快也不嫌快先说一句很多人可能不太在意的背景知识。深度学习模型跑一次推理计算量绝大部分落在卷积和矩阵乘法上而卷积在底层实现时也往往会转换成矩阵乘法来处理。所以 GEMM 的性能直接决定了一次前向推理的快慢。GEMM 的数学公式非常简单C alpha * A * B beta * C。A 是 M×K 矩阵B 是 K×N 矩阵C 是 M×N 输出矩阵。看起来就是三重循环的事但问题在于数据量远超芯片的片上存储。一颗现代 GPU 有几百 GB/s 甚至几 TB/s 的显存带宽但一个 4096×4096 的 FP16 矩阵就有 32MB整个模型矩阵动辄几十个这样的块不可能全部塞进片上缓存。所以 GEMM 优化的核心从来不是怎么算而是怎么搬数据搬得少、算得密。我选择自己写 DeepGEMM不是因为官方闭源库不行而是因为它解决不了我的三个问题。第一融合算子。推理引擎里 GEMM 后面基本都跟着 bias、激活函数、LayerNorm、量化缩放这些操作官方库只负责输出 C 矩阵剩下的操作要我重新起一个 kernel 读写一遍数据显存带宽全浪费了。第二形状控制。官方库对超大连续矩阵调校得很好但碰到 M1 的推理场景、或者 N 比较小的窄矩阵性能优势就没那么明显了。第三可控性。我想针对自己的模型形状做定制想在看到性能瓶颈时能精确知道是哪一行代码在等待闭源库做不到。1.2 DeepGEMM 的定位不是重新造轮子是造一套能改的轮子DeepGEMM 的定位是一套面向深度学习推理场景的 GEMM 内核模板集不是要替代官方库在所有场合的工作而是重点覆盖推理引擎里最常用到的几种情况常见批量大小下的大矩阵乘、低精度输入FP16/BF16甚至量化后的 INT8/FP8、以及需要把额外算子焊进 epilogue 的融合需求。做这套东西之前我先明确了一个原则先让内核在一个小形状上跑出明显的性能再考虑泛化而不是一开始就试图处理所有边界情况。我最初只实现了 M4096、N4096、K4096 的基础版本跑通之后再逐步往外扩。这样做的原因是 GEMM 优化的变量实在太多tile 尺寸、寄存器布局、流水线深度、访存模式任何一个改动都可能改变性能特征如果不固定形状去调参很难知道到底是哪个改动起了作用。2. 分块调度把大矩阵切成能塞进芯片的小方块2.1 从数学公式到三层存储层级矩阵乘法最直观的实现就是三重循环累加但这样每个 A 元素会被读 N 次每个 B 元素会被读 M 次数据搬运量惊人。分块的目标是提高数据复用A 矩阵的某一行会被 C 矩阵同一行的所有列使用B 矩阵的某一列会被 C 矩阵同一列的所有行使用所以把计算切成 M_BLOCK × N_BLOCK 的小方块后这个小方块计算只需要加载对应的 M_BLOCK×K 的 A 分块和 K×N_BLOCK 的 B 分块数据复用率从 1 提升到了块尺寸级别。在 GPU 上分块要分层进行。第一层是把整个 C 矩阵分成若干 M_BLOCK × N_BLOCK 的块每个 block线程块负责一个输出分块第二层是 block 内部把 K 维再切段每次从显存加载小块 A、B 到共享内存第三层是每个线程从共享内存取数据用寄存器算对应的输出小片累加结果最终写回显存。这个三层结构恰好对应 GPU 的三种存储层级显存、共享内存、寄存器。2.2 一个实例算清楚分块参数怎么定举一个具体例子。假设目标 GPU 有 132 个 SM目标形状是 MNK4096。我最初选的 block 尺寸是 128×128那 grid 就是 32×32 1024 个 block平均每个 SM 要处理约 7.8 个 block负载基本均衡。K 维切片长度 BLOCK_K 的选择更讲究。BLOCK_K 越大单次加载的数据越多访问显存效率更高但共享内存占用也随之增加。128×128 的 C 分块用 FP32 累加寄存器需要 128×128/32 线程 每线程 64 个 FP32 寄存器加上操作数寄存器寄存器压力已经不小。所以 BLOCK_K 我选了 32这样 A 分块是 128×32×2 字节 8KBB 分块是 32×128×2 字节 8KB加上 C 分块共享内存占用大概 20KB 左右在 128KB 的共享内存里可以轻松放下双层缓冲。提示BLOCK_K 一旦超过 64共享内存占用会急剧上升留给双层缓冲的空间就紧张了。实际测试下来在多数数据中心级 GPU 上BLOCK_K 在 32 到 64 之间是最稳的选择区间。选好这些参数之后主循环的控制流就非常机械了外层遍历 K 维切片内层把切到的 A、B 分块从显存搬到共享内存然后所有线程执行矩阵乘的小片计算。问题是这种搬一块、算一块的做法搬数据的时候计算单元是空闲的计算的时候搬运单元是空闲的性能只有理论峰值的一半左右。解法就是第 3 节要说的双层缓冲。3. 真正提速的核心细节张量核指令、寄存器排布、流水线预取3.1 mma 指令与 FP32 累加为什么低精度输入要用高精度加法现代数据中心级 GPU 架构都有一个专门做矩阵乘法的硬件单元通常称为张量核心它通过底层的 mma 指令一次性完成一个小矩阵的乘累加。比如一条 mma 指令可以完成 16×8×16 这样的操作也就是 A 是 16×16、B 是 16×8算出 16×8 的输出小矩阵。向量单元要一个时钟周期做几次乘加张量核心一个周期能完成一个片段的乘累加吞吐量完全不在一个量级。DeepGEMM 里我采用的指令模式是加载 FP16/BF16 输入累加器用 FP32。这不是我拍脑袋定的——直接用 FP16 累加会在多次累加后出现明显的舍入误差尤其在 K 很大时误差会累积。FP32 累加寄存器虽然占的地方多一倍但换来的是结果精度明显提高实测中与官方库的 FP32 累加结果完全相同。PTX 层的指令看起来大概是这样mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32 {%0, %1, %2, %3}, // 输出累加器 16x8 的四个片段 {%4, %5}, // A 矩阵片段 16x16 {%6, %7}, // B 矩阵片段 16x8 {%0, %1, %2, %3}; // 输入累加器这条指令的难点在于操作数布局是固定的A、B、C 的每个片段分别放在哪些寄存器、哪些位都有严格规定。如果寄存器排布不符合指令预期编译器会补一堆 mov 指令来回倒腾数据性能直接打折。所以我写 DeepGEMM 时是先按指令要求的寄存器布局来分配数据而不是先分配好再期望编译器去适配。3.2 共享内存 bank 冲突看不见的串行化陷阱共享内存按 bank 分块硬件在同一周期可以同时服务多个不同 bank 的访问。但如果多个线程访问的是同一个 bank 的不同地址吞吐就会变成原来的几分之一这就是所谓 bank conflict。在 GEMM 内核里,加载 B 矩阵分块时最常踩这个坑。经典案例以 FP16 类型为例BLOCK_N 为 128B 分块是 BLOCK_K × 128。如果每个线程连续取一行中的相邻元素两个线程的地址恰好落在同一个 bank 上整个 load 就会被串行化。解决方法是给 B 矩阵的访问模式加一个偏置移位让相邻线程访问的地址错开到不同 bank。DeepGEMM 里我用了一个简单的 swizzle 策略矩阵存储时同一行的数据不再连续排列而是按 XOR 变换重新组织。这样线程在列方向并行读取时地址映射后的 bank 号自然错开实测共享内存带宽利用率从约 60% 提升到接近 90%。以 C 语言伪码来说明一个 8 元素 swizzle 的思路// 每个线程要读取的元素所在行 row、列 col // 传统布局address row * row_pitch col // swizzle 布局address row * row_pitch \ // ((col 7) ^ ((row 1) * 7))关键是让相邻线程的列地址和行本身的偏移撞不到同一个 bank。3.3 双层缓冲让搬运和计算真正重叠软件流水线的目标是让显存到共享内存的拷贝和矩阵乘计算重叠执行。最朴素的做法是共享内存里放两份 A、B 分块一份给当前循环迭代使用另一份给下一次迭代预取。主循环写成这样而循环里的操作步骤是用 cp.async 指令异步发起下一次迭代的 A、B 数据拷贝用当前共享内存缓冲区里的数据执行 mma 指令完成矩阵乘提交计算完成等待之前的异步拷贝完成交换两个缓冲区的角色进入下一轮cp.async 是 GPU 上的一种异步拷贝指令数据从显存搬进共享内存不占用线程计算时间硬件会自动完成。第一次实现时我直接把 cp.async 换成普通 load性能下降接近两成原因就是每次主循环结束都得等数据搬运完才能继续算。这个经验在 DeepGEMM 的开发中反复被验证只要共享内存装得下双缓冲几乎是零成本提升。4. 精度对齐与排错链路从错误结果到逐项排查4.1 第一次跑通不代表结果是错的我第一版 DeepGEMM 跑出了结果乍一看和官方库的输出差不多但差分对不过最大绝对误差在 1e-1 量级对推理模型来说这直接是错误输出。当时我最怀疑的是张量核指令用错了后来发现其实是边界处理的问题。排查的第一原则是固定变量。我把输入矩阵改成随机数固定种子先用官方库算出基准 C_ref再用 DeepGEMM 算出 C_test逐项对照。拿 M、N、K 都比较小的用例开始比如 64×64×64这样任何一行代码的行为都能手动推算。4.2 排查流程从索引、同步到边界我当时整理的排查顺序现在也推荐给你检查项方法我遇到的问题索引公式在小矩阵上逐元素核对 A、B 分块加载确认 mma 指令的 A 矩阵行主序转列主序时搞混了一次同步位置检查 __syncthreads() 是否覆盖跨线程共享数据读写有一次预取数据已经发起当前计算还没用完旧缓冲区就做了缓冲区交换边界处理M、N 不能被 block 整除时越界位置是否置 0128 的 block 处理 4096 没问题换到 1000 就出错了累加精度与 FP32 基准对比误差是否在可接受区间FP16 累加误差在 K4096 时放大到不可接受4.3 两个最容易踩的坑越界加载与缓冲区交换先说越界加载。BLOCK 尺寸通常是 32 或 64 的倍数但矩阵 M、N 不一定是整数倍。当尾部块加载 A 分块时最后几行已经超出矩阵边界读到的显存内容是什么不确定算出的结果自然不对。解法是加载时做边界判断越界位置置零。这个判断放在主循环里太贵我把它放在每次加载的 if 分支里if (row M col K) { A_tile[local_row * BLOCK_K local_col] A[global_row * K global_col]; } else { A_tile[local_row * BLOCK_K local_col] 0.f; // 越界置零不影响累加结果 }置零的好处是无论该位置被计算多少次对最终累加结果都没有贡献尾块照样可以走和完整块相同的 mma 指令。再说缓冲区交换。双缓冲实现里有一个非常隐蔽的 bug我在异步拷贝还没完成时就把缓冲区的指针换掉了导致下次计算用的还是旧数据。排查这个问题的代码路径是用__pipeline_commit和__pipeline_wait_prior这对异步接口时忘记在读取共享内存前等待对应批次全部完成。注意多级流水线里等待条件不是上一批拷贝完成而是我这次计算需要的那一批完成。流水线越深这个关系越容易搞错。我后来用了一个比较实用的检验方法把 K 维的切片数改成奇数跑一遍如果结果和偶数切片不一致基本就是流水线同步有 bug。因为切片数量奇偶变化会改变缓冲区和循环迭代的对应关系任何等待顺序错误都会导致结果异常。5. 性能实测用 profiler 数据驱动不拍脑袋调参5.1 基线对比要分形状不能只看一个数GEMM 性能受形状影响极大。我拿 DeepGEMM 和官方闭源库做了对比测试固定数据格式为 FP16 输入 FP32 累加分别测了三种典型形状结果很能说明问题形状M × N × KDeepGEMM 达到的峰值算力占比与官方库相对性能4096 × 4096 × 4096约 78%约 95%1024 × 1024 × 4096约 70%约 88%1 × 4096 × 4096约 15%约 20%前两个结果说明 DeepGEMM 在常规大矩阵上已经有竞争力最后一个 M1 的形状则暴露了问题block 里大量线程计算同一行数据复用不够内存加载变成了瓶颈。这个结果提醒我DeepGEMM 目前的架构并不适合下沉到 M 很小的推理场景还需要针对这种情况做独立优化。5.2 用 profiler 找到瓶颈而不是靠猜性能出问题的时候我先用 GPU 的性能分析工具采集三个指标SM 忙碌率、共享内存带宽利用率、请求停滞分布。第一次跑 4096 形状SM 忙碌率只有 55%共享内存带宽利用率也不是特别高说明问题不在计算而在流水线等待。逐个排查后发现主循环里每轮计算结束后有一个等待异步拷贝的操作按我的预期这里应该已提前预取完毕。但分析工具显示 stall 主要发生在共享内存访问阶段原因是加载 B 矩阵分块时的 swizzle 只做了部分偏移仍有少量 bank 冲突。我把 swizzle 粒度从 4 元素调整到 8 元素之后共享内存带宽利用率从 72% 上升到了 88%整体算力占比也提升了大约 8 个百分点。5.3 编译器参数同样影响性能写 GEMM 内核时编译参数的影响经常被低估。我用到的两个关键选项一个是指定 GPU 架构让编译器针对具体指令集优化另一个是开启更大的寄存器限制允许寄存器溢出到局部的优化。最开始我用默认编译参数寄存器分配比较保守性能差了约一成。把寄存器上限调高到 255 之后编译器能把更多的中间结果留在寄存器里而不是反复写回共享内存。需要注意的是这里有个平衡点寄存器占用太高会导致 SM 上同时运行的线程块变少并行度下降。我在目标 GPU 上实际测试128 线程一 block、每个线程 168 个寄存器左右是可以同时跑两个 block 的上限区间。6. 往融合算子走一步GEMM 的 epilogue 定制6.1 推理场景真正需要的是 GEMM 偏置 激活如果 DeepGEMM 只能输出一个 C 矩阵那它对推理引擎的价值就打折了。实际上全连接层之后的模式非常固定先加 bias再过 ReLU/GELU最后才是输出。把这三个步骤从独立算子合并进 GEMM kernel 的末尾段epilogue可以减少一次完整的显存读写。我在实现上把 epilogue 做成一个模板化的函数传入输出块的共享内存指针和模型配置由每个线程在算完自己的输出片段后执行。模板参数里包含是否需要 bias、使用哪个激活函数、是否要做量化这些全都在编译期确定运行时零分支开销。6.2 量化场景把 scale 和 zero point 焊进融合流程低精度推理里GEMM 输出是 FP32但下一层要求 INT8 输入中间必然有一次量化操作y round(x * scale zero_point)。常规做法是把 C 矩阵写回显存再启动一个量化 kernel 读出来重新算一遍带宽和时间成本都很高。在 DeepGEMM 的 epilogue 中我在每个线程输出 FP32 片段之后直接做这个量化再把结果写回显存。如果量化是 per-token 的即每行一个 scale也只需要把 scale 向量提前加载到共享内存每个线程按自己的行索引取数即可。融合后一次 kernel 跑完 GEMM 和量化实测整体的显存读写量下降约一半端到端时间缩短了约 30%。6.3 融合的边界与后续扩展方向融合也不是越多越好。层归一化虽然也在 GEMM 之后常用但它需要跨通道计算均值方差意味着要拿到一行的完整输出而 GEMM 的输出片段是按 block 分散在不同线程里的强行融合会导致复杂的跨线程规约。我在 DeepGEMM 里没有把 LayerNorm 合进去而是让它保持独立 kernel这也是很多推理引擎的普遍做法。这套内核后续可以扩展的方向我目前最关心的是更小的 block 尺寸以适配推理批次较小的场景以及把 FP8 输入的计算路径补全因为在最新一代硬件上 FP8 的性能收益非常明显。另一个思路是为结构化稀疏做配套优化稀疏矩阵在带宽上的节省潜力比纯稠密更大。我实际写 DeepGEMM 的体会是矩阵乘法的性能优化没有什么玄学所有瓶颈最后都能落到访存模式、寄存器分配、流水线等待这几个具体问题上。关键是不要一开始就追求大而全而是从一个小形状出发把 profiler 给出的数据一项项磨平。你先跑通一个 block 都行把同步、边界、精度都验证对了再往上加复杂度和融合特性这个过程会稳很多。
返回列表
PREV
查看更多资讯
NEXT
返回资讯列表