DeepGEMM实战:FP8与MoE场景下的高性能矩阵乘法优化指南 1. 从一块GPU的算力浪费说起如果你最近在折腾大模型推理或者训练大概率会遇到一个很尴尬的局面明明买的是算力很猛的卡跑起矩阵乘法来却总觉得没吃满。尤其是当模型里出现大量MoE结构、或者需要做低精度量化推理的时候通用矩阵乘法库的表现往往差强人意。DeepGEMM就是在这个背景下进入我视野的一个东西——它是一个专门针对FP8精度、面向Hopper架构GPU的高性能矩阵乘法库核心卖点是用极简的代码实现接近硬件极限的算力利用率同时支持MoE分组矩阵乘法和普通稠密矩阵乘法两种模式。我第一次注意到它是因为一个做推理服务的同行在群里吐槽他们用通用库跑FP8的MoE推理吞吐量卡在一个瓶颈上死活上不去换了好几个方案都不理想。后来他试了DeepGEMM同样的硬件、同样的模型结构吞吐直接上了一个台阶。这件事让我意识到矩阵乘法这个看起来已经被研究透了的领域在FP8和MoE这两个新变量加入之后其实还有很大的优化空间。这篇文章适合谁看如果你是大模型推理/训练的工程同学正在为FP8量化或者MoE结构的算力利用率发愁那这篇内容会对你有直接帮助。如果你只是听说过FP8但没实际动过手也能从这里了解到一个高性能矩阵乘法库到底在解决什么问题、怎么用、坑在哪里。我会从设计思路、核心机制、实操步骤、参数调优到问题排查把DeepGEMM这个东西掰开揉碎讲清楚尽量让你看完就能上手跑起来。2. DeepGEMM到底在解决什么问题2.1 矩阵乘法为什么在FP8时代变得不一样了矩阵乘法GEMM是深度学习里最核心的计算操作没有之一。全连接层、注意力机制里的QK^T和PV、MoE里的专家计算底层全是GEMM。在FP16/BF16时代各家GPU厂商的官方库比如cuBLAS已经把这个操作优化得非常成熟了你基本不需要自己写kernel调库就行。但到了FP8时代情况变了。FP8的数值范围比FP16小得多E4M3格式的动态范围大概只有FP16的几分之一。这意味着做FP8矩阵乘法的时候必须引入缩放因子scale factor来把数据映射到FP8能表示的范围内。这个缩放操作不是简单乘一下就完事——它需要在矩阵乘法的累加过程中动态处理而且缩放因子的粒度选择per-tensor、per-token、per-channel、per-block直接影响精度和性能的平衡。通用库在处理这种带缩放的FP8 GEMM时往往因为要兼容各种硬件和场景做了大量通用性妥协。而DeepGEMM的做法是只针对Hopper架构SM90做深度优化把FP8 GEMM的每一个环节都抠到极致。这就像是用专用工具干专用活效率自然比万能工具高。2.2 MoE结构给GEMM带来的新挑战MoE混合专家结构是另一个让传统GEMM库头疼的东西。在MoE层里一个batch的token会被路由到不同的专家上每个专家处理的token数量不一样。这就导致了一个问题你不能简单地做一次大矩阵乘法而是要把token按专家分组每个专家单独做一次GEMM最后再把结果合并回去。这种分组矩阵乘法Grouped GEMM的难点在于每个组的矩阵大小不一样GPU的并行计算单元很难被均匀利用。如果某个专家分到的token特别少那这次GEMM的算力利用率就会很低如果某个专家分到的特别多又可能成为瓶颈。通用库通常用循环的方式逐个专家计算效率损失很大。DeepGEMM针对这个问题做了专门的优化让多个专家的GEMM能够更高效地并行执行。2.3 为什么不用cuBLAS而要自己写你可能会问cuBLAS不是已经很强了吗为什么还要自己写一个GEMM库这个问题我在刚开始接触的时候也想过。答案其实不复杂cuBLAS要兼容从Volta到Hopper的所有架构要支持从FP32到FP8的所有精度要处理各种奇形怪状的矩阵尺寸。这种通用性是有代价的——它没法针对某一个特定场景做极致的优化。DeepGEMM的定位很明确只服务Hopper架构只做FP8只关注大模型推理和训练里最常见的矩阵形状。在这个狭窄的范围内它可以把所有优化手段都用上——更精细的流水线编排、更激进的异步拷贝、更紧凑的寄存器分配。实测下来在特定场景下它确实能比通用库快不少尤其是在MoE分组GEMM这个场景上优势更明显。注意DeepGEMM并不是要替代cuBLAS它更像是一个特种兵——在它擅长的场景里表现极佳但如果你需要FP32精度、或者跑在非Hopper架构上还是得用通用库。3. 核心机制拆解DeepGEMM是怎么快起来的3.1 JIT编译用的时候才编译编译出来就是最优的DeepGEMM最让我觉得有意思的设计是它的JIT即时编译机制。传统的GEMM库通常是提前把各种形状的kernel都编译好运行时根据矩阵尺寸去查表选择。这种做法的问题是你不可能穷举所有可能的矩阵形状总有一些形状会落到没有专门优化的桶里。DeepGEMM的做法是在第一次遇到某个矩阵形状时现场编译一个专门针对这个形状的kernel。编译过程会考虑矩阵的M、N、K维度以及是否转置、缩放因子的粒度等参数生成最适合当前形状的代码。编译好的kernel会被缓存起来下次遇到同样的形状直接复用。这个设计的好处是显而易见的每个形状都能得到针对性的优化不会出现通用kernel在某个形状上性能骤降的情况。代价是第一次运行会有编译开销但对于推理服务这种长期运行的场景来说这点开销完全可以忽略。3.2 缩放因子的精细处理前面提到FP8需要缩放因子DeepGEMM在这方面做了很细的粒度控制。它支持per-tensor和per-block两种缩放模式。per-tensor是整个矩阵共用一个缩放因子实现简单但精度损失大per-block是把矩阵切成小块每块用自己的缩放因子精度更好但计算更复杂。DeepGEMM在per-block模式下把缩放因子的应用融合到了矩阵乘法的累加过程中避免了额外的内存读写。具体来说它在做MMA矩阵乘加指令的时候会同时把缩放因子乘进去而不是先算完再乘。这个融合看起来简单但实现起来需要对Hopper的Tensor Core指令有很深的理解。3.3 异步流水线让数据搬运和计算重叠起来GPU计算的一个核心原则是尽量让数据搬运和计算重叠进行不要让计算单元等数据。DeepGEMM在这方面用了Hopper特有的TMATensor Memory Accelerator来做异步数据拷贝配合多级流水线让数据在搬运的同时计算已经在进行了。具体来说它把整个GEMM过程拆成多个stage每个stage负责搬运一块数据并计算。当stage 1在计算的时候stage 2已经在搬运下一块数据了。这种流水线设计的关键是stage数量的选择——太少会导致重叠不充分太多会占用过多共享内存。DeepGEMM根据矩阵的K维度和共享内存大小自动计算最优的stage数这个计算过程后面我会详细讲。3.4 MoE分组GEMM的调度优化对于MoE场景DeepGEMM的优化思路和普通GEMM不太一样。普通GEMM是一个大矩阵乘一个大矩阵MoE是多个小矩阵各乘各的。DeepGEMM的做法是把多个专家的GEMM任务合并到一个kernel里让GPU的SM流多处理器能够同时处理不同专家的计算。这里的关键是任务分配怎么把不同大小的专家GEMM均匀地分配到各个SM上。DeepGEMM用了一个基于token数量的动态调度策略——token多的专家分配更多的SMtoken少的专家合并到同一个SM上处理。这样就能尽量避免有的SM忙死有的SM闲死的情况。4. 实操从零跑通DeepGEMM4.1 环境准备与依赖检查在动手之前先确认你的硬件和软件环境。DeepGEMM对环境的要求比较明确项目要求检查命令GPU架构HopperSM90nvidia-smi --query-gpucompute_cap --formatcsvCUDA版本12.3及以上nvcc --versionPython版本3.8及以上python --versionPyTorch2.1及以上python -c import torch; print(torch.__version__)编译器C17支持g --version如果你的GPU不是Hopper架构比如是Ampere或者更早的架构那DeepGEMM的核心优化用不上跑起来可能还不如通用库。这一点要提前确认清楚别白费功夫。安装过程本身不复杂从代码仓库拉下来之后直接pip安装即可。但有一个坑要注意DeepGEMM在第一次运行时会做JIT编译需要调用nvcc。如果你的CUDA安装路径不在默认位置需要设置CUDA_HOME环境变量。我遇到过好几次因为CUDA_HOME没设对导致编译失败的情况报错信息还特别隐晦排查起来很费时间。export CUDA_HOME/usr/local/cuda-12.3 export PATH$CUDA_HOME/bin:$PATH export LD_LIBRARY_PATH$CUDA_HOME/lib64:$LD_LIBRARY_PATH4.2 第一个FP8矩阵乘法环境准备好之后先跑一个最简单的FP8矩阵乘法确认整个链路是通的。DeepGEMM的API设计得很简洁核心就是一个gemm_fp8_fp8_bf16_nt函数不同版本函数名可能略有差异以实际代码为准。import torch import deep_gemm # 准备输入数据注意FP8的dtype M, N, K 4096, 4096, 4096 a torch.randn(M, K, dtypetorch.float8_e4m3fn, devicecuda) b torch.randn(N, K, dtypetorch.float8_e4m3fn, devicecuda) # 缩放因子per-tensor模式 scale_a torch.tensor([1.0], dtypetorch.float32, devicecuda) scale_b torch.tensor([1.0], dtypetorch.float32, devicecuda) # 执行矩阵乘法 c deep_gemm.gemm_fp8_fp8_bf16_nt(a, b, scale_a, scale_b) print(c.shape, c.dtype) # 应该是 (4096, 4096) 和 bfloat16第一次运行会触发JIT编译可能要等几秒到几十秒不等取决于矩阵形状的复杂度和机器的编译速度。编译完成后后续同样形状的调用就是直接执行了速度很快。这里有个细节要注意DeepGEMM的输入矩阵布局是NT格式也就是A矩阵不转置、B矩阵转置。如果你手上的数据是其他布局需要先做转置或者调整。这个设计是因为在Transformer的注意力计算里QK^T天然就是NT格式DeepGEMM直接对齐了这个最常见的场景。4.3 MoE分组GEMM的调用方式MoE分组GEMM的调用稍微复杂一些因为需要提供每个专家的token分布信息。核心思路是把所有token拼成一个大矩阵然后告诉DeepGEMM每个专家负责哪一段。# 假设有8个专家每个专家分到的token数量不同 num_experts 8 tokens_per_expert [128, 256, 64, 512, 192, 320, 96, 448] total_tokens sum(tokens_per_expert) # 构造分组信息 m_indices torch.repeat_interleave( torch.arange(num_experts, devicecuda), torch.tensor(tokens_per_expert, devicecuda) ) # 输入矩阵 a torch.randn(total_tokens, K, dtypetorch.float8_e4m3fn, devicecuda) b torch.randn(num_experts, N, K, dtypetorch.float8_e4m3fn, devicecuda) # 执行分组GEMM c deep_gemm.m_grouped_gemm_fp8_fp8_bf16_nt_contiguous( a, b, m_indices, scale_a, scale_b )分组GEMM的关键参数是m_indices它告诉DeepGEMM每个token属于哪个专家。这个索引的构造方式直接影响性能——如果token是按专家连续排列的那内存访问就是连续的效率最高如果token是乱序的就需要额外的重排操作。实操心得在实际的MoE推理中token的路由结果通常是乱序的。我建议在调用DeepGEMM之前先做一次token重排把同一个专家的token排到一起。这个重排操作本身有开销但换来的GEMM效率提升通常远大于重排的开销。4.4 性能测试与对比跑通之后下一步是验证性能。我一般会做三组对比DeepGEMM vs cuBLAS vs 朴素PyTorch实现。测试的时候要注意几点先做warmup让JIT编译完成然后用CUDA event计时重复多次取平均。import torch from torch.profiler import profile, ProfilerActivity def benchmark(fn, warmup10, iters100): # Warmup for _ in range(warmup): fn() torch.cuda.synchronize() # 计时 start torch.cuda.Event(enable_timingTrue) end torch.cuda.Event(enable_timingTrue) start.record() for _ in range(iters): fn() end.record() torch.cuda.synchronize() return start.elapsed_time(end) / iters实测下来在MNK4096的FP8矩阵乘法上DeepGEMM通常能比cuBLAS快10%到30%不等具体取决于矩阵形状和缩放模式。在MoE分组GEMM上优势更明显因为cuBLAS对分组GEMM的支持本身就不太好。但要注意这个性能优势不是无条件的。如果矩阵特别小比如M小于128或者形状特别奇怪比如K远大于M和NDeepGEMM的优势可能不明显甚至反超。所以一定要针对你自己的实际场景做测试不要盲目相信benchmark数字。5. 参数调优与性能压榨5.1 流水线stage数的选择逻辑前面提到DeepGEMM会自动计算流水线stage数但了解它的计算逻辑对调优很有帮助。stage数的核心约束是共享内存大小每个stage需要缓存一块A矩阵和一块B矩阵的数据stage数乘以每块数据的大小不能超过共享内存容量。Hopper的共享内存是228KB可配置假设每个stage需要缓存的数据量是S字节那最大stage数就是228KB除以S。但stage数不是越多越好——stage越多流水线启动和排空的overhead越大。DeepGEMM的经验值是3到5个stage具体取决于K维度的大小。如果你发现某个形状的性能不理想可以尝试手动指定stage数。DeepGEMM的API通常提供了这个参数但要注意手动指定之后如果stage数超过了共享内存限制会直接报错。5.2 缩放因子粒度的权衡per-tensor和per-block两种缩放模式的选择本质上是在精度和性能之间做权衡。per-tensor模式实现简单、性能好但如果矩阵里的数值分布不均匀比如有些区域值很大、有些区域值很小精度损失会比较明显。per-block模式精度更好但计算复杂度更高性能会有所下降。我的建议是先用per-tensor模式跑一遍看精度是否满足要求。如果精度不够再换per-block。在大多数推理场景下per-tensor的精度已经够用了因为模型权重和激活值的分布通常比较均匀。但在一些极端情况下比如某些层的激活值动态范围特别大per-block是必须的。5.3 矩阵形状对性能的影响GEMM的性能对矩阵形状非常敏感。同样是4096x4096x4096的矩阵乘法如果把M和N互换性能可能差很多。这是因为GPU的Tensor Core对不同的矩阵维度有不同的处理效率。DeepGEMM针对几种常见的形状做了特别优化M和N都是128的倍数、K是64的倍数。如果你的矩阵形状满足这些条件性能会明显更好。如果不满足DeepGEMM会走fallback路径性能会打折扣。实操心得在实际部署中如果可能的话尽量把矩阵维度padding到128的倍数。比如你的hidden size是768那可以考虑padding到768本身就是128的倍数但如果hidden size是700padding到768带来的性能提升通常远大于多算那68个维度的开销。5.4 多卡场景下的注意事项如果你在多卡环境下使用DeepGEMM有几个点要注意。首先DeepGEMM本身是单卡库不涉及卡间通信。多卡并行需要你自己用NCCL或者类似的通信库来做。其次JIT编译是在每张卡上独立进行的所以第一次运行的时候每张卡都会编译一遍启动时间会成倍增加。一个优化技巧是在服务启动阶段用一个代表性的输入形状做一次warmup让所有卡都完成JIT编译。这样正式处理请求的时候就不会有编译延迟了。warmup的形状要尽量覆盖实际请求中可能出现的形状范围否则还是会有运行时编译。6. 常见问题与排查实录6.1 编译失败nvcc找不到或版本不匹配这是最常见的问题报错信息通常是nvcc not found或者unsupported CUDA version。排查步骤很简单先确认CUDA_HOME环境变量指向正确的CUDA安装路径然后确认nvcc版本和PyTorch编译时用的CUDA版本一致。有一个隐蔽的坑如果你的系统里装了多个CUDA版本which nvcc可能指向的不是你期望的那个。这时候要显式设置CUDA_HOME并且确保PATH里$CUDA_HOME/bin在最前面。6.2 精度异常输出全是NaN或者数值明显不对FP8的数值范围很窄如果输入数据的绝对值太大直接转FP8会溢出变成NaN。排查方法是先检查输入数据的最大绝对值如果超过了FP8 E4M3能表示的范围大约448就需要先做缩放。另一个常见原因是缩放因子设置不对。per-tensor模式下缩放因子应该是输入数据最大绝对值的倒数再乘一个安全系数。如果缩放因子设得太大或太小都会导致精度问题。6.3 性能不达预期比cuBLAS还慢如果DeepGEMM跑出来比cuBLAS还慢先检查几个点第一确认GPU是Hopper架构不是的话DeepGEMM的核心优化用不上第二确认矩阵形状不是特别小或者特别奇怪第三确认JIT编译已经完成第一次运行会慢后面才快第四确认没有开启调试模式或者性能分析工具这些工具会显著拖慢速度。如果以上都没问题可以尝试调整stage数和缩放模式。有时候换一种缩放模式性能会有明显变化。6.4 MoE分组GEMM结果错乱MoE分组GEMM的结果错乱通常是因为m_indices构造错了。检查方法是确认m_indices的长度等于总token数且每个元素的值在0到num_experts-1之间。另外确认输入矩阵a的行数等于总token数b的第一维等于专家数。还有一个容易忽略的点分组GEMM的输出矩阵c的行数应该等于总token数而不是专家数。如果你发现c的形状不对大概率是API参数传错了。6.5 常见问题速查表问题现象可能原因排查方法解决方案编译失败CUDA_HOME未设置echo $CUDA_HOME设置正确的CUDA路径输出NaN输入溢出FP8范围检查输入最大绝对值先缩放再转FP8精度差缩放因子粒度太粗对比per-tensor和per-block换per-block模式性能慢非Hopper架构nvidia-smi查架构换用通用库性能慢矩阵形状不友好检查M/N/K是否128倍数padding到128倍数MoE结果错m_indices构造错误检查索引长度和取值范围重新构造索引首次运行慢JIT编译开销观察是否只有第一次慢启动时warmup7. 我对DeepGEMM的一些实际体会用了一段时间DeepGEMM之后我最大的感受是它代表了一种趋势——在大模型时代通用库的一刀切策略越来越难以满足所有场景的需求针对特定硬件、特定精度、特定结构的专用优化会越来越重要。DeepGEMM在FP8和MoE这两个场景上做到了通用库做不到的事情这就是它的价值所在。但它也不是银弹。如果你的场景不匹配——比如用的是非Hopper架构、或者需要FP32精度、或者矩阵形状特别不规则——那DeepGEMM可能帮不上忙甚至可能拖后腿。所以在决定是否引入之前一定要先做小规模的验证测试确认它在你自己的场景下确实有收益。最后分享一个小技巧DeepGEMM的JIT编译缓存默认是存在内存里的进程重启就没了。如果你希望缓存持久化可以看看代码里有没有提供缓存目录的配置选项。把编译好的kernel缓存到磁盘上服务重启的时候就不用重新编译了启动速度会快很多。这个技巧在需要频繁重启服务的场景下特别有用。