ARTICLE · INTELLIGENCE

战地情报 · 详情页

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

CUTLASS实战指南:从GEMM源码到AI推理性能优化

CUTLASS实战指南:从GEMM源码到AI推理性能优化 CUTLASS这个词最近在GPU高性能计算和AI推理优化圈子里出现频率越来越高凡是想把手里的模型压榨到极致的人几乎都绕不开它。它不是什么玄乎的算法库而是NVIDIA开源的一套基于CUDA C模板头文件库专门用来实现高性能矩阵乘GEMM、卷积、Attention等深度学习核心算子而且它把整个底层设计思路摊开给你看从线程块到线程再到张量核每一层的任务划分、数据搬运、计算调度都可以由你自己用代码定义。这篇文章我不会讲那种复制粘贴的API文档而是从源码结构、工程能力、CuTe这套DSL的底层原理到最后怎么在真实AI推理场景里落地下手给你一条完整的路线。适合想做算子优化、推理引擎开发、或者单纯想搞懂NVIDIA底层计算库怎么工作的人读完之后你能大致建立对CUTLASS整体设计和实际使用的感觉。1. 先搞清楚CUTLASS到底是什么不只是又一个GEMM库1.1 一句话版本与适用人群CUTLASS的全称是CUDA Templates for Linear Algebra Subroutines直译过来就是“用于线性代数子程序的CUDA模板库”。它最核心的定位是提供高性能的通用矩阵乘法GEMM模板。你可能会问NVIDIA不是已经有cuBLAS和cuDNN了吗为什么还要造一个CUTLASS因为cuBLAS、cuDNN这类闭源库是黑盒你没法改内部计算策略比如你想给矩阵乘融合一个量化缩放、一个ReLU、一个自定义的注意力掩码cuBLAS做不到cuDNN的融合也有限你只能把数据来回搬运好几次性能损耗巨大。CUTLASS就给了你完全可控的路径你自己定义数据从哪里读、用什么方式写入共享内存、用哪条张量核指令、后处理阶段做什么融合。所以CUTLASS的适用人群很明确做推理引擎的开发工程师、做算子库的团队、研究GPU架构性能的学术人员、以及想在Jetson等边缘设备上榨干每一分算力的嵌入式开发者。如果你只是调用现成的PyTorch模型跑推理那大概率不需要直接碰CUTLASS但如果你用到的推理框架底层发生了性能瓶颈而你又想搞清楚瓶颈在哪、怎么改CUTLASS能给你答案。1.2 为什么NVIDIA要自己开源CUTLASS很多人在第一次接触CUTLASS时都会有疑问NVIDIA不靠卖CUDA软件赚钱吗把高性能GEMM的核心实现直接开源不怕别人学了去其实NVIDIA开源的逻辑很清晰GPU硬件越卖越多但要让用户真正发挥出硬件性能光靠闭源库是不够的。NVIDIA的客户里有一大批是做AI框架、高性能计算中间件的这些人需要深度定制算子如果只有黑盒库他们要么自己做不满性能要么去逆向硬件指令最终反而拖慢生态。开源CUTLASS相当于给出了一个“官方最佳实践模板”告诉大家怎么写才能达到接近硬件上限的性能。同时CUTLASS里的代码结构、模板设计、调度思想也对下一代硬件设计有反馈意义——比如Hopper架构上的TMA引擎、张量内存加速器这些新特性都需要有对应的软件抽象来体现价值而CUTLASS就是这套抽象的试验场。另外还有一点容易被忽略CUTLASS是很多GPU性能分析的标杆。不管是做算子性能调优还是做编译器代码生成都需要一个“和顶尖手写实现对比”的基线CUTLASS就扮演了这个角色。我自己在做算子优化时经常先跑一遍CUTLASS的profiler看看同规模下CUTLASS跑到多少TFLOPs心里就有底了——如果我自己写的核函数连CUTLASS的80%都到不了那先别谈什么更高级的优化先老老实实把内存访问模式对齐。2. 架构拆解从线程块到张量核的完整链路2.1 三层并行体系与数据流CUTLASS的性能之所以强核心在于它把GPU的并行模型理解得非常透彻并且用模板参数把每一层的决策都暴露出来。一个典型的GEMM核函数内部计算是被层层划分的。最外层是线程块CTA级别的tile比如把整个输出矩阵切成128x128或者128x256的块每个线程块负责一个输出子块。线程块内部又分成多个warp级的tile每个warp再分成线程级的tile最后每个线程通过张量核指令MMA一次算出若干个浮点结果。整个数据流是典型的“全局内存 - 共享内存 - 寄存器 - 张量核”的多级搬运路径。为什么要这么折腾因为GPU的全局内存延迟非常高可能有几百个时钟周期而共享内存带宽高、延迟低很多寄存器最快。CUTLASS的做法是先从全局内存把一块A矩阵的数据搬进共享内存再从共享内存搬运到寄存器喂给tensor core。在这个搬运算的过程中计算和搬运是重叠进行的也就是主循环mainloop里每个迭代都同时做“下一块数据的拷贝”和“当前块数据的矩阵乘”用软件流水线把访存延迟藏掉。要注意的是这一步是CUTLASS设计的灵魂。很多人自己写CUDA GEMM写得慢就是因为在等数据乘法器大部分时间处于空闲状态。CUTLASS通过ping-pong双缓冲甚至多缓冲让共享内存始终有下一块数据等着被消费计算单元几乎停不下来。这也是为什么CUTLASS在Ampere、Hopper等新架构上能跑到理论峰值80%以上的原因之一。2.2 主循环与后处理算得快还要写得对CUTLASS把GEMM核函数在逻辑上拆成两个大阶段mainloop主循环和epilogue尾声/后处理。Mainloop阶段就是不断执行“搬运A、B数据 累加计算”这两个动作把C矩阵的所有累加结果留在寄存器中。Epilogue阶段则是在所有累加完成后把寄存器中的结果写出到全局内存并在这个过程里融合各种操作线性变换alpha/beta缩放、激活函数ReLU、GELU等、量化反量化、甚至逐元素加上偏置。这个拆分的价值在于解耦。你在做算子融合时只需要修改epilogue部分的模板参数把想要的融合操作以函数对象方式传进去而不需要动主循环的复杂流水线。比如做注意力算子可以让输出之前叠加一个softmax或者mask这在传统cuBLAS里是根本做不到的。我自己在实际项目里最常用的就是把FP32累加结果转成INT8输出顺便把per-channel的scale和zero point乘进去这一个融合操作就能省下两到三次全内存访问整个层耗时降了30%。后处理阶段还有一个关键细节是数据的存储布局。如果你的输出需要转置、需要按channel-padded方式存储epilogue阶段就要用不同的shared memory加载和全局写入模式。CUTLASS允许你通过模板参数指定输出布局这让它能够适配非常多框架的需求而不需要你在框架层去额外做一次内存重排。2.3 算子族全景不止是FP16矩阵乘很多人以为CUTLASS只做FP16的GEMM这是大误解。它的算子覆盖范围相当广从数据类型看有FP16、BF16、TF32、FP32、INT8、INT4以及Hopper和Blackwell上的FP8从计算形态看有标准的GEMM、split-K GEMM把K维切分到多个CTA分别计算再归约、群组GEMMGrouped GEMM一次处理多个不同形状的小矩阵、流式GEMMstream-K还有基于GEMM派生出来的卷积实现以及近年来大模型必需的Flash Attention实现。Grouped GEMM在推荐系统和LLM推理里非常关键因为真实场景中往往同时有多个不同长度的序列需要计算如果逐个调用标准GEMM切换开销很大。CUTLASS的grouped GEMM允许你把一组不同M维度的GEMM打包进一次kernel launch。还有stream-K这种模式它把K维的累加工作拆分到更多CTA上然后对局部累加结果进行二次归约在矩阵规模不够大、并行度不足的情况下stream-K能明显提升GPU占用率我在处理M只有几百的全量推理层时靠stream-K拿到了差不多20%的性能提升。3. CuTe DSL原理把张量布局变成可编程的语言3.1 CuTe的核心抽象张量与布局CuTe是CUTLASS 3.x引入的一套底层抽象库全称是CUDA Templates for Linear Algebra很多人把它理解为一种嵌入式DSL。为什么需要CuTe因为CUTLASS早期的版本里布局、切分、内存映射这些逻辑都是散落在各个类里的代码里满屏的sizes、strides、swizzle函数新人根本看不懂而且想要重新组合一种新的数据流模式要改动的模板代码非常多。CuTe要做的就是把这些概念提炼成几个基本积木Layout布局、Tensor张量视图、Tile切分方式、Copy拷贝操作、MMA矩阵乘操作然后让这些积木可以彼此组合。这里的核心概念是Layout。一个Layout描述了张量到内存的映射关系可以用(shape, stride)的二元组表示。形状是各个维度的大小步长指示每个维度上的元素在内存中相隔多远。在CuTe中你会看到类似LayoutShape_4, _8, Stride_8, _1这种写法它表示一个4行8列的矩阵行方向步长为8、列方向步长为1也就是标准的行优先二维矩阵。但Layout的威力远不止于此它还能表达任意复杂的映射比如把int8数据按4字节打包的vectorized布局、用于消除共享内存bank conflict的swizzle布局等。3.2 布局代数形状、步长与复合CuTe的DSL感来自它的一套“布局代数”规则。你可以对两个Layout做乘积、做复合、做逆、做拼接。比如make_layout(A, B)表示把两个布局拼接成一个更大的布局composition(L, R)表示先按R布局切分数据再按L布局组织这些数据这正好用来定义tile在原始矩阵上的嵌套映射。还有一个极其常用的操作是logical_product和tiled_product它们用来生成多层tile组合整个线程块级别的tile、线程束级别的tile、线程级别的tile全都是用这些代数操作组合出来的。把布局抽象成代数对象之后最大的好处是代码可以泛化。你写一个计算逻辑不需要关心矩阵到底是行优先、列优先还是被swizzle过、被pad过只要传入对应的LayoutCuTe能够在编译期做大量的常量折叠与推导自动知道每个线程该从哪个地址取数。这就好像你在Python里写一个对列表做变换的算法不管底层列表怎么存储都能跑只要符合迭代协议就行。CuTe把这种协议定义得很严密而且代价为零因为所有信息都在编译期完成解析。我自己第一次看CuTe的头文件时最震撼的是它大量使用了整数常量模板形状里的维度全是_4、_8这种编译期常量而不是运行时的int。这意味着所有循环边界在编译期就确定了循环可以完全展开地址偏移可以立即算出常量生成的SASS非常干净。代价就是模板实例数量爆炸编译时间很长后面我会专门讲这个坑。3.3 在流水线里使用CuTe的实例举一个非常具体的CuTe使用场景。假设你想实现一个F32 GEMM的线程级切分线程块负责128x128的输出tile每个线程计算8x8的结果。在CuTe里你要做的是用make_tile(cute::Int128{}, cute::Int128{})定义CTA级别的tile定义线程布局通常是一个8x8的线程块组织用logical_product将CTA tile和线程布局组合得到每个线程负责的元素坐标再对A矩阵做对应的partition得到每个线程的A、B、C视图在mainloop里调用CUTE_TMA_LOAD或者cute::copy完成数据搬运调用cute::gemm执行MMA累加。这样的代码天然适配不同的架构指令。在Ampere上cute::gemm可能展开成mma.m16n8k8的PTX指令在Hopper上会变成新的wgmma指令。你不用去记ARCH的差异只需要把Layouts定义好CuTe会自动做匹配。这就是它DSL价值的最大体现你描述的是“数据应该怎么组织”而不是每一步硬件指令的具体寄存器编号。不过要注意CuTe上手曲线是有的尤其是从常规CUDA C转过来的开发者会不太适应这种“一切皆模板”的风格。我建议循序渐进先跑通CUTLASS官方自带的CuTe教程比如cutlass/tools/util中的layout例程把Layout打印出来看理解(_8,_8):(_8,_1)到底是一张什么样的内存map再上手写自己的kernel会顺畅很多。4. CUTLASS的工程能力从代码生成到核函数融合4.1 编译期多态与代码实例化CUTLASS的一大工程特点是“尽量在编译期做决策”。所有关键参数——tile大小、线程数、指令类型、流水线级数、融合操作——都以模板参数的方式传入编译器会为每一组参数生成一份独立的CUDA kernel代码。这套机制带来两个结果。好的一面是零运行时开销。当你针对特定形状和特定硬件实例化一个kernel时NVIDIA编译器nvcc可以把大量模板元编程计算在编译期完成循环展开、寄存器分配、shared memory地址计算全都提前做好最后生成的机器码可以直接逼近手写优化汇编的水平。坏的一面是编译时间极其夸张模板稍微组合多一点编译一个带十几个kernel的算子库动辄要几十分钟而且最终二进制体积会非常大。我见过一个同事只是改了模板的一个参数重启编译后去喝了杯咖啡回来还没编完。应对这个问题的工程经验是分层组织代码。把运行时会变化的字段比如M、N的具体数值映射到常量模板参数时要克制不要每个维度都开一个独立的实例。合理的做法是提供一个通用的运行时形状dispatch函数再配合典型的tile配置集合比如64x64、128x64、128x128做几个预编译实例运行时根据形状选择一个最接近的。这样既能保持性能又不会让编译时间失控。4.2 算子融合与内存规划CUTLASS工程能力最亮眼的部分我觉得是它对算子融合的系统化支持。前面提到的epilogue融合只是一种。在CUTLASS 3.x中你还可以通过collective builder把更复杂的融合编排进去比如多阶段的global-shared拷贝、TMA的异步搬运、多缓冲的流水线调度。算子融合的意义在AI推理里非常直观模型中的每个算子都在消费和产生张量如果这个张量只能放在全局内存里每多一次算子就需要多两轮全内存访问写一次、读一次。在内存带宽远低于算力的时代这个开销让硬件只能发挥三四成算力。CUTLASS允许你把“矩阵乘缩放激活残差”一次性算完中间的结果全部躺在寄存器或者共享内存里根本不落盘。对一个GPT类模型来说单层Attention中融合掉几个中间张量端到端的延迟能明显下降。内存规划上CUTLASS也让开发者可以直接控制共享内存的分配与复用。你可以为多个操作分配同一块共享内存只要它们的生命周期不重叠。由于共享内存只有几十到上百KB精细的复用规划往往能决定一个kernel能不能跑起来以及能不能用上更高的流水线级数。说实话这块非常考验工程经验同一个GEMM不同共享内存规划方案性能可以差到两倍以上。4.3 评测与基准测试注意点CUTLASS源码仓库自带了一套性能测试工具可以针对不同输入规模跑GEMM基准输出TFLOPS和带宽利用率这是评价你调优成果最直接的手段。但有几个测试细节要提醒一下一是GPU的时钟频率问题跑benchmark前最好用nvidia-smi -lgc锁定GPU的base clock和memory clock不然测试过程中频率波动会导致数据不可比。二是要预热前几次kernel调用会把CUDA context建立、cuModuleLoad之类的开销戴进来跑出来的数字会偏低要循环跑多轮取稳定值。三是矩阵对齐条件CUTLASS性能对指针地址对齐非常敏感常见的128位向量化load要求16字节对齐如果数据指针没有特殊对齐性能可能会掉一半以上。另外做基准测试时不要只看峰值TFLOPs还要看实际带宽受限度。很多GEMM在规模较小时其实是带宽瓶颈而不是算力瓶颈此时再优化计算tile的尺寸也没有大用应该考虑减少数据搬运或者提高数据复用率。我对团队新人的建议是先跑一个cuBLAS同规模下的GEMM作为baseline确保CUTLASS能持平或略优于cuBLAS再依次调整tile大小、流水线级数、swizzle函数每一步都记录性能变化这样才不会在优化中迷失方向。5. AI推理落地指南从Demo到生产环境的路径5.1 与PyTorch等框架集成的几种方式CUTLASS作为一个底层库当然不会直接被PyTorch调用但它有两种主流的集成路径。第一种是直接用CUTLASS的Python扩展接口CUTLASS Python即cutlass模块它提供了类似PyTorch扩展的绑定可以把你定义的CUTLASS kernel封装成PyTorch的torch.autograd.Function然后就能直接在nn.Module里用。这种方式非常适合快速验证算子融合效果我最初做Layernorm融合GEMM的实验就是通过这个路径跑通的。第二种更底层的方式是直接编写CUDA C扩展在C层调用CUTLASS的template类然后用pybind11或TorchScript的custom op接入框架。这种方式灵活度最高也能把CUTLASS的核函数嵌套在更大的自定义逻辑里面但它要求你对CUTLASS API、内存管理、以及PyTorch的autograd机制都比较熟调试成本更高。如果做的是纯推理还有一条更高效的路径跳过PyTorch直接用CUTLASS构建一个轻量推理引擎。把离线转换好的权重按CUTLASS的最优布局提前转换好例如把权重在H维度上做K-major连续存储把量化参数打包进同一块缓存。这样在服务端推理时省掉了权重转换和内存拷贝用CUTLASS直接跑实测在INT8 GEMM场景里能达到比cuBLAS还高的吞吐这在LLM长上下文解码阶段尤为关键。5.2 在真实推理引擎中扮演的角色现在主流的推理引擎和加速方案比如TensorRT、vLLM、FasterTransformer现在是TensorRT-LLM、乃至一些大厂的内部推理平台其底层高性能算子在很大程度上都是参考或直接复用CUTLASS的。有些引擎甚至直接内置了CUTLASS的kernel。我理解CUTLASS在推理引擎里的角色主要有两个。一个是作为“兜底高性能实现”引擎对用户不常用的算子不可能像主推算子那样做极致定制但又要保证不能比cuBLAS慢太多CUTLASS作为通用高性能GEMM模板正好填这个坑。另一个角色是作为“新架构适配先行者”每当NVIDIA出一代新硬件比如BlackwellCUTLASS往往最早适配新的张量核指令和新的内存流水线第三方引擎会先看CUTLASS是怎么用新特性的再去跟进实现自己的专用算子。在实际部署LLM时要特别注意推理的两个阶段对GEMM形状的巨大差异。Prefill阶段是典型的compute-bound大GEMMM等于输入序列长度可能很大应该用CUTLASS里为大规模规整GEMM优化的kernel。而Decode阶段一次只生成一个tokenM非常小这时更像GEMV矩阵向量乘瓶颈在带宽而不是算力用split-K或者stream-K去提高并行度才是正解。选错kernel形态性能会差非常多这是我踩过的实实在在的坑。5.3 Jetson等边缘设备上的实战建议CUTLASS不仅能在数据中心GPU上跑它在NVIDIA Jetson系列比如Jetson Orin NX/AGX上同样可用这对边缘AI推理非常有价值。Jetson使用的是集成式GPU架构上通常与Ampere或上一代Turing/Volta类似CUDA核心数量、共享内存大小、张量核能力都和桌面显卡差距很大因此同一个CUTLASS kernel在Jetson上的最佳tile配置跟在A100上完全不同。在Jetson上做算子部署我建议重点检查几件事。第一是CUDA toolkit版本要和JetPack SDK一起管理不要单独装桌面版CUDA否则nvidia-smi和驱动版本对不上最常见的就是“nvidia-smi has failed because it couldnt communicate with the nvidia driver”这类报错处理方法通常是重新安装匹配的驱动或者重启后重新加载内核模块。第二是Jetson的GPU内存和CPU共享同一块物理内存数据搬运路径和独立显卡不一样用CUDA统一内存可能比显式拷贝还要高效这时可以适当减少shared memory的stage数量节省共享内存给更深的软件流水线。第三是要考虑交叉编译如果是在x86主机上编译再拷贝到Jetson跑要确保指定了aarch64的架构常用的做法是用JetPack自带的交叉编译工具链或者直接在板上用gcc和nvcc原生编译后者虽然慢一点但省心很多。边缘部署还要注意功耗限制。Jetson默认有功率墙TDP限制如果算子为了追求极致性能把CUDA core全部跑满功耗会迅速冲到上限反而触发降频得不偿失。实际工程里往往要在性能与功耗之间找平衡比如把tile调小一点减少并发线程数让GPU在较低频率下稳定工作往往端到端吞吐反而更高。6. 常见问题与排查技巧实录6.1 编译期模板地狱CUTLASS最劝退新人的就是报错信息经常是几百行模板实例化错误叠在一起看起来完全无从下手。我的经验是做“最小复现”先把完整kernel拆掉单独编译一个最简单的矩阵乘例子确认环境无误后再逐步叠加自己的修改。很多时候编译错误出在类型不匹配上比如一个表达shape的常量类型写成了int而不是cute::Int这类错误在模板库里很常见但阅读错误信息时一定要看最底层的那几行“required from here”那才是真正的问题位置。还有一个实用技巧是充分使用C的static_assert。在你自己的kernel代码里加入对维度、对齐条件、共享内存大小的静态断言这样一旦模板参数不合法你能立刻看到自己写的明确提示而不是被一堆晦涩的模板堆栈淹没。另外把编译任务拆小当你一次实例化了太多kernel组合编译出问题定位就会很困难可以先只编译一个tile配置验证通过后再放开所有配置。6.2 性能不达标的排查思路当你发现CUTLASS kernel跑出来的性能远低于预期时不要急着改模板参数先按顺序做三件事。第一件是检查Profiler数据用Nsight Compute跑一遍看是访存带宽受限、计算受限还是延迟受限。如果显示“Memory Throughput”很高但“Compute Throughput”很低说明瓶颈在数据搬运需要增加共享内存stage或者调整swizzle来减少共享内存bank冲突反过来如果计算占用率已经很高说明可能是tile太大导致占用率不足并行warp太少。第二件是检查L2 cache命中率和内存访问是否合并。全局内存访问如果不对齐或者不连续CUTLASS的高效load模式会失效性能损失巨大。一个常见问题是我前面提到的16字节对齐使用cutlass::fast_math或者reinterpret_cast时很容易忽略这一点导致向量化load退化成非向量化逐元素load。第三件是用Nsight Systems做一次跨kernel的分析看kernel与kernel之间的间隙、GPU空转时间有时候性能问题不在kernel内部而是kernel launch之间CPU端任务调度太慢这是API overhead问题要减少launch次数或使用CUDA Graph做捕获与回放。6.3 环境与工具链的几个坑最后说一些环境方面容易踩的坑这些坑我几乎每次换新机器都会遇到一次。第一个是驱动、CUDA Toolkit、CUTLASS三者的版本要匹配。CUTLASS的新特性比如TMA、FP8依赖较新的CUDA版本和驱动如果你用的是老驱动即使代码逻辑正确也可能在编译期或运行期卡住。装驱动时不要图省事直接用系统源里的nvidia驱动包建议用NVIDIA官网匹配显卡型号的runfile或deb包装完一定要确认nvidia-smi能正常输出否则后续任何CUDA计算都跑不起来。第二个坑是gcc版本不匹配比如Ubuntu高版本默认的gcc可能太新导致nvcc编译报错这时候需要安装一个CUDA Toolkit官方支持的gcc版本并且通过PATH或CXX环境变量指定。第三个是显存和共享内存的检查在调试时可以用cuda-memcheck或compute-sanitizer跑一遍能快速定位非法内存访问但注意这类sanitizer工具会显著降低运行速度只适合在debug阶段用。最后一个建议是给关键代码加详细的注释尤其是layout推导的部分。CUTLASS代码里有一大堆编译期值不看注释半年后连自己都看不懂。我现在养成的习惯是每个tile配置都写清楚为什么选这个shape、针对什么带宽/算力场景、实测性能数据是多少。这样后续做回归对比或者别人接手项目时才不会需要重新踩一遍我之前踩过的所有坑。我自己在项目中的体会是CUTLASS不是那种看一遍文档就能拿来就用的库第一次接触至少需要一周时间通读官方示例才能算真正入门。但一旦你会熟练拆分和组合它的模板组件那种感觉就像从只能呼叫黑盒API变成自己掌握了一台发动机的每一个零件。遇到性能问题时你不再需要去猜黑盒里发生了什么而是可以直接打开Nsight看每一条流水线改一个swizzle函数再看效果这种调试链路是cuBLAS永远给不了你的。如果你正在做AI推理落地并且碰到算力瓶颈强烈建议静下心来把CUTLASS的核心源码读一遍你会对“性能”这两个字有完全不同的理解。
RELATED READING

延伸阅读

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