CUDA 算子工程:手写 FlashAttention v2 之路
第 13 章 CUTLASS 3.x 设计哲学
CUTLASS 与其说是一个要"学会"的库,不如说是一套要"用起来"的词汇:CollectiveOp、CuTe layout、TileScheduler。 讲清楚这套词汇,是为了让读者能看出 FA3、Machete 这些 Hopper kernel 共用的是同一副骨架。
13.1 CUTLASS 是什么
CUTLASS(CUDA Templates for Linear Algebra Subroutines)是 NVIDIA 官方维护的开源 C++ 模板库,第一版 2017 年 Volta 时代发布,到现在(本专栏钉的 v4.7.0)已经成熟到:
- 体量极大:v4.7.0 里
include/一个目录就有近 69 万行头文件(CUTLASS 是 header-only 的模板库,没有src/;算上examples/、test/、tools/与 Python 侧,整个仓库超过 180 万行)。 - 覆盖 Volta / Turing / Ampere / Hopper / Blackwell 五代架构。
- 支持几十种数据类型组合(FP16 / BF16 / TF32 / FP8 / INT8 等)。注意 INT4 只存在于 Turing / Ampere 的 mma.sync 通路上——
cutlass-4.7.0/include/cutlass/arch/mma_sm75.h与cutlass-4.7.0/include/cutlass/arch/mma_sm80.h里各有 24 处int4b_t,而 Hopper 的 WGMMA 封装cutlass-4.7.0/include/cute/arch/mma_sm90_gmma.hpp里整数只有 s8/u8 这一档,没有 s4。INT4 的mma.sync在 sm_90a 上虽然还能编译,但 nvcc 13.4 生成的 SASS 是先用整数指令把 INT4 解包成 INT8,再发两条IMMA.16832.S8.S8;同一段代码编 sm_80 则是一条原生的IMMA.16864.S4.S4(不同 CUDA 版本可能不同)。也就是说,Hopper 的 Tensor Core 没有原生 INT4 MMA。第 9 章讲的 INT4 权重量化,在 Hopper 上是"存成 INT4、算之前先解回 FP16/BF16",不是"用 INT4 算"。 - 被大量开源 kernel 直接采用:FlashAttention v3(
flash-attn/hopper/)、vLLM 的 Machete 与vllm-0.8.5/csrc/quantization/cutlass_w8a8/、PyTorch 的分组 GEMM(pytorch-v2.11.0/aten/src/ATen/native/cuda/GroupMM.cu)等。cuBLAS 闭源,NVIDIA 官方文档的说法是 CUTLASS 采用的分层分解与数据搬运策略与实现 cuBLAS 所用的类似(cutlass-4.7.0/media/docs/cpp/doxygen_mainpage.md:7)。
简单说:今天开源世界里的高性能 LLM GEMM 算子,很多直接建在 CUTLASS 上。
但 CUTLASS 也是出名的难学。几十万行模板代码、深度嵌套的类型层级、上百个 traits class——新人打开 CUTLASS 源码常常会陷入"看 5 分钟模板就晕"的状态。
这一章的目标不是教读者用 CUTLASS(它的 API 还在演化),而是教读者理解 CUTLASS 的设计哲学——CuTe Layout、CollectiveOp 三段式、Kernel Schedule。理解这些之后,读者再打开 CUTLASS 源码会发现"哦原来这里在做这件事"。
13.2 CUTLASS 三代演进
CUTLASS 的设计经过了三次大重构,每一次都是对硬件能力的重新抽象:
flowchart LR
subgraph V1 [CUTLASS 1.x · 2017-2019]
V1A[Volta · WMMA 16×16×16]
V1B[平铺 GEMM 模板]
V1C[Bottom-up: 用底层指令组合]
end
subgraph V2 [CUTLASS 2.x · 2019-2022]
V2A[Turing/Ampere · mma.sync 16×8×8 / 16×8×16]
V2B[Threadblock-level GEMM 抽象]
V2C[Iterator + Pipeline 模式]
end
subgraph V3 [CUTLASS 3.x · 2023+]
V3A[Hopper · WGMMA + TMA]
V3B[CuTe Layout + CollectiveOp]
V3C[Top-down: 描述意图,自动展开]
end
V1 --> V2 --> V3
(各版日期见 cutlass-4.7.0/CHANGELOG.md:0.0.1 是 2017-12,2.0.0 是 2019-11,3.0.0 是 2023-01。4.x 在 C++ 之外又加了 Python 的 CuTe DSL,与 CuTe C++ 是同一套抽象(cutlass-4.7.0/README.md:30);本章只讲 C++ 这一侧。)
13.2.1 1.x 时代:模板地狱
CUTLASS 1.x 是 Volta 架构的产物。它的核心抽象是"GEMM Pipeline"——把 GEMM 拆成 prologue / mainloop / epilogue 三段,每段都是模板类。
代码风格大致这样:
template <
typename ElementA, typename LayoutA,
typename ElementB, typename LayoutB,
typename ElementC, typename LayoutC,
int ThreadBlockShapeM, int ThreadBlockShapeN, int ThreadBlockShapeK,
int WarpShapeM, int WarpShapeN, int WarpShapeK,
int MmaShapeM, int MmaShapeN, int MmaShapeK,
typename EpilogueOp
>
class Gemm { ... };
模板参数动辄十几个。能写出来,但巨难看懂、巨难改。
13.2.2 2.x 时代:Iterator 与 Pipeline
CUTLASS 2.x 引入了几个关键抽象:
- Iterator:把"从某个 Tensor 中以某种 stride/layout 读出 fragment"的过程模板化。Iterator 隐藏了地址计算和 vectorized load。
- Pipeline:把 prologue + mainloop + epilogue 用模板组合起来,自动处理 double buffer。
代码可读性大幅提升,但还是有大量隐式约定("layout 必须满足 contract X"),新人难入门。
13.2.3 3.x 时代:CuTe + CollectiveOp
CUTLASS 3.x 是真正的范式重构。两个核心抽象:
- CuTe:一个独立的 header-only 子库(
include/cute/),专门做 layout 代数。把"什么样的数据怎么放"用一种通用语言描述出来。 - CollectiveOp:把 GEMM 拆成 CollectiveMainloop(核心循环,类名是
cutlass::gemm::collective::CollectiveMma)和 CollectiveEpilogue(输出阶段),每段是一个高层抽象类。
3.x 的代码大致长这样:
using CollectiveMainloop = cutlass::gemm::collective::CollectiveBuilder<
cutlass::arch::Sm90, cutlass::arch::OpClassTensorOp,
half_t, LayoutA, 8,
half_t, LayoutB, 8,
float,
Shape<_128, _128, _32>,
Shape<_1, _1, _1>, // ClusterShape
cutlass::gemm::collective::StageCountAuto,
cutlass::gemm::KernelTmaWarpSpecialized
>::CollectiveOp;
模板参数还是多,但每个参数的含义清晰,而且 stage 数、kernel schedule、epilogue tile 这些不好拍脑袋定的参数都有 Auto 选项,交给 builder 推导。这就是 3.x 的进步——降低使用门槛,但保留自定义能力。
13.3 CuTe Layout:CUTLASS 3.x 的灵魂
CuTe 是 CUTLASS 3.x 中所有概念的基础。它的核心是把"如何排布张量"形式化。
13.3.1 Layout 的定义
CuTe Layout 是 (Shape, Stride) 的 pair:
using L = Layout<Shape<_4, _8>, Stride<_1, _4>>; // 4×8 矩阵, column-major (M-major)
Shape (4, 8) 表示 4 行 8 列。Stride (1, 4) 表示行下标 i 加 1 偏移加 1、列下标 j 加 1 偏移加 4——元素 (i, j) 落在 i*1 + j*4,也就是列优先(column-major)。行优先要写成 Stride<_8, _1>:
矩阵地址: 元素 (i, j) 的偏移:
(0,0) -> 0 (0,1) -> 4 ... offset = i * 1 + j * 4
(1,0) -> 1 (1,1) -> 5 ...
这是基本的 layout。CuTe 强大在于能组合 layout:
// 嵌套 layout: 4×8 矩阵被分成 2×4 个 (2×2) 子块, 每个子块内部 row-major、占连续 4 个偏移,
// 子块之间也按行排: 行 (i0, i1) 的 stride 是 (2, 16), 列 (j0, j1) 的 stride 是 (1, 4)
using L = Layout<
Shape<Shape<_2, _2>, Shape<_2, _4>>,
Stride<Stride<_2, _16>, Stride<_1, _4>>
>;
用 print_layout 打出来,第 0 行是 0 1 4 5 8 9 12 13、第 1 行是 2 3 6 7 10 11 14 15:左上角 2×2 子块正好占偏移 0~3。复杂吗?是的。但 Tensor Core fragment 布局、线程到数据的映射这类不规则布局,都能用 (Shape, Stride) 嵌套表示;swizzle 则是在这之上再复合一个 XOR 函数(13.3.3 节)。
13.3.2 Layout 代数
CuTe 提供一组对 Layout 的操作:
composition(A, B):函数复合 A∘B,先用 B 映射坐标、再交给 A(cutlass-4.7.0/include/cute/layout.hpp:1136)。logical_divide(A, B)及其变体zipped_divide/tiled_divide:把 A 按 B 切成 tile(同文件 :1559、:1610、:1621)。logical_product(A, B)及其变体blocked_product/raked_product:把 A 这个 block 按 B 的排布复制铺开(同文件 :1653、:1734、:1752)。这是 layout 的乘积,不是笛卡尔积。complement(A, M):求 A 在大小为 M 的空间里的"补",divide 和 product 都靠它实现(同文件 :1234)。layout(tensor):取张量的 layout(cutlass-4.7.0/include/cute/tensor_impl.hpp:525)。
这些操作让你可以像写代数公式一样描述 GEMM 的数据流:
// 摘自 cutlass-4.7.0/examples/cute/tutorial/sgemm_1.cu:103-121 与 :199,只留 A 这一路
auto cta_coord = make_coord(blockIdx.x, blockIdx.y, _); // (m,n,k)
Tensor gA = local_tile(mA, cta_tiler, cta_coord, Step<_1, X,_1>{}); // (BLK_M,BLK_K,k) 本 block 要读的 A
Tensor sA = make_tensor(make_smem_ptr(smemA), sA_layout); // (BLK_M,BLK_K) SMEM 里的 tile, layout 自选
Tensor tAgA = local_partition(gA, tA, threadIdx.x); // (THR_M,THR_K,k) 按线程布局 tA 切出本线程那份
Tensor tAsA = local_partition(sA, tA, threadIdx.x); // (THR_M,THR_K)
// 主循环里:
copy(tAgA(_,_,k_tile), tAsA); // 本线程那份 gmem -> smem
copy(src, dst) 的默认重载会根据两边的 layout 自动决定向量宽度(cutlass-4.7.0/include/cute/algorithm/copy.hpp:307):两边 layout 全是编译期常量时,按 128 位对齐的假设把能合并的元素合并成宽访存;只要有一边含运行时 stride(比如 stride 运行时才知道的全局内存 tensor),就只按 8 位对齐来假设,half 这类 16 位元素于是逐个搬(想要宽访存可以改用同文件的 copy_aligned)。要换成 cp.async 或 TMA,则是显式换一个 Copy Atom:传 AutoCopyAsync 策略(同文件 :173,gmem→smem 时选 SM80_CP_ASYNC_*),或者用 make_tma_copy(cutlass-4.7.0/include/cute/atom/copy_traits_sm90_tma.hpp:1332)建 TMA 描述符、再配上 mbarrier。循环本身不用改,程序员也不需要手写 cp.async 指令。
这就是 CuTe 最革命性的地方——用一种声明式语言描述数据流,把"用哪条指令搬、用哪条指令算"收拢到 Copy Atom / MMA Atom 这一个选择点上。
13.3.3 Swizzle 在 CuTe 中
第 12 章讲的 swizzle layout 在 CuTe 中是一个一等公民:
using SwizzleAtom = decltype(
composition(Swizzle<3, 3, 3>{}, Layout<Shape<_8, Int<BK>>, Stride<Int<BK>, _1>>{}) // BK = 64 个 half
);
Swizzle 的三个模板参数顺序是 Swizzle<BBits, MBase, SShift>(cutlass-4.7.0/include/cute/swizzle.hpp:55 里分别叫 num_bits / num_base / num_shft):MBase 是"不打散的基本单元"占多少位,BBits 是拿多少位去做 XOR,SShift 是这个 XOR 源从哪一段位取。别把顺序记反。GMMA 那一族 *_Atom_Bits 用的是 Swizzle<0,4,3> / <1,4,3> / <2,4,3> / <3,4,3>(cutlass-4.7.0/include/cute/atom/mma_traits_sm90_gmma.hpp:75-:84),MBase = 4 是因为那一族带 smem_ptr_flag(cutlass-4.7.0/include/cute/pointer_flagged.hpp:53),swizzle 作用在 SMEM 指针的字节地址上(cutlass-4.7.0/include/cute/pointer_swizzle.hpp:87 把指针转成 uintptr_t 再做 XOR), 字节 = 16 字节一块,Swizzle<3,4,3> 就是用字节地址的第 7–9 位去 XOR 第 4–6 位,形成 128B 模式;名字里的 _Bits 指的是 layout 部分以比特为单位(_1024 比特 = 一行 128 字节),按元素类型 upcast 之后 swizzle 仍是 Swizzle<3,4,3>(用 CuTe 的 print 打印 Layout_K_SW128_Atom<half_t>,得到的是 Sw<3,4,3> o smem_ptr[16b] o (_8,_64):(_64,_1))。同一档 128B swizzle 写成直接作用在元素下标上的 layout 就是 Swizzle<3,3,3>( 个 half = 16 字节),FlashAttention 的 kernel_traits.h 用的正是后一种写法,第 15 章 §15.4.4 会对上。CUTLASS 在 SM90::GMMA 命名空间里把它们包成了几组预定义的 layout atom(Layout_MN_SW128_Atom<T> / Layout_K_SW128_Atom<T> 等,同文件 :87 起),覆盖 WGMMA 支持的几档 swizzle。
13.4 CollectiveMainloop:核心循环的抽象
CollectiveMainloop 描述 GEMM 的核心循环——从 HBM 拉 A/B tile 到 SMEM,从 SMEM 读 fragment 到寄存器,调用 mma 累加。
CUTLASS 3.x 提供了一组 Kernel Schedule:
// cutlass-4.7.0/include/cutlass/gemm/dispatch_policy.hpp
KernelMultistage // :112 Ampere, multi-stage cp.async pipeline
KernelTma // :117 Hopper, TMA 但不做 warp specialization
KernelTmaWarpSpecialized // :118 Hopper, TMA + Warp Specialized
KernelTmaWarpSpecializedPingpong // :119 Hopper, 两个 MMA warpgroup 交替占用 Tensor Core
KernelTmaWarpSpecializedCooperative // :122 Hopper, 两个 MMA warpgroup 协作算同一个 tile
(2.x 时代的经典双缓冲不在这张表里——它属于旧的 cutlass::gemm::threadblock::MmaPipelined 那一套,3.x 的 dispatch policy 没有对应项。)
每个 Schedule 对应一个 kernel 层实现(cutlass-4.7.0/include/cutlass/gemm/dispatch_policy.hpp:110 的注释原话是 "one for each kernel layer file",对应 cutlass-4.7.0/include/cutlass/gemm/kernel/ 下的 sm90_gemm_tma_warpspecialized*.hpp 等),同一个 mainloop 可以配不同的 schedule(cutlass-4.7.0/media/docs/cpp/gemm_api_3x.md:311)。后三种都实现了第 2 章讲的 Producer/Consumer warp specialization;KernelTmaWarpSpecialized 是其中最朴素的一种:1 个 producer warp group 加 1 个 MMA warp group(cutlass-4.7.0/include/cutlass/gemm/kernel/sm90_gemm_tma_warpspecialized.hpp:140-:141),不做持久化,一个 tile 发一个 CTA(同文件 :254 起的 get_grid_shape)。实际选型时,KernelScheduleAuto 在 CUDA 12.1 及以上会选持久化的 Cooperative(tile M ≥ 128)或 Pingpong(tile M = 64),源码注释说持久化 schedule 在这些版本上表现最好(cutlass-4.7.0/include/cutlass/gemm/collective/builders/sm90_gmma_builder.inl:1018-:1025)。
伪代码:
// 内部展开后的 KernelTmaWarpSpecialized mainloop
__global__ void mainloop() {
if (warp_group == 0) {
// Producer warp group: 持续发起 TMA
for (int k = 0; k < K_tiles; ++k) {
// 等这个 stage 被 consumer 释放, 并在 full barrier 上 arrive_and_expect_tx 登记要到的字节数
producer_acquire(k % stages);
cp_async_bulk_tensor(smem_a[k % stages], tma_desc_a, k, full_barrier[k % stages]);
cp_async_bulk_tensor(smem_b[k % stages], tma_desc_b, k, full_barrier[k % stages]);
// 不再单独 arrive: TMA 搬到的字节自己在 full barrier 上记账, 凑齐即翻转相位
}
} else {
// Consumer (MMA) warp group: 等数据并算 WGMMA
for (int k = 0; k < K_tiles; ++k) {
mbarrier_wait(full_barrier[k % stages]);
wgmma_mma_async(c_acc, smem_a[k % stages], smem_b[k % stages]);
wgmma_commit_group();
wgmma_wait_group(0); // 简化: 真实实现是 wait<1>, 让一组 wgmma 留在飞行中, 释放的是上一个 stage
release_pipeline(k % stages);
}
}
}
读者不需要手写这段代码——KernelTmaWarpSpecialized 把它实现好了。读者只要在 CollectiveBuilder 模板参数里指定就行。
13.5 CollectiveEpilogue:输出阶段的抽象
GEMM 的 epilogue 是 D = activation(alpha * acc + beta * C + bias) 这种最终阶段。CUTLASS 3.x 把它抽象成 CollectiveEpilogue:
using CollectiveEpilogue = cutlass::epilogue::collective::CollectiveBuilder<
cutlass::arch::Sm90, cutlass::arch::OpClassTensorOp,
Shape<_128, _128, _32>,
Shape<_1, _1, _1>,
cutlass::epilogue::collective::EpilogueTileAuto,
float, float,
half_t, LayoutC, 8,
half_t, LayoutD, 8,
cutlass::epilogue::TmaWarpSpecialized, // 与 mainloop 的 KernelTmaWarpSpecialized 配对
cutlass::epilogue::fusion::LinearCombination<half_t, float>
>::CollectiveOp;
epilogue 的 schedule 要和 mainloop 的 kernel schedule 配对:KernelTmaWarpSpecialized 配 TmaWarpSpecialized(cutlass-4.7.0/examples/61_hopper_gemm_with_topk_and_softmax/61_hopper_gemm_with_topk_and_softmax.cu:124-:125),KernelTmaWarpSpecializedCooperative 配 TmaWarpSpecializedCooperative。
最后一个参数 LinearCombination 是融合操作。可以替换成自定义 epilogue:
// fusion 命名空间下的写法: 线性组合 + 逐元素激活, 激活函数作为模板参数传进去
using FusedAddRelu = cutlass::epilogue::fusion::LinCombEltAct<
cutlass::epilogue::thread::ReLu, half_t, float>;
cutlass-4.7.0/include/cutlass/epilogue/fusion/operations.hpp 里列着这一族:LinearCombination(:116)、LinCombEltAct(:131,任意逐元素激活)、LinCombPerRowBias(:161,加 bias)、LinCombPerRowBiasEltAct(:196,bias 加激活一起)、LinCombEltActBlockScaleFactor(:536,激活之外顺带生成输出的块级 scale factor,给块缩放的低精度输出用)等等。旧的 cutlass::epilogue::thread::LinearCombinationRelu 这类"一种融合一个类"的写法仍在(2.x 遗产),但 3.x 的方向是把激活函数抽成模板参数。如果这些组合还不够用,可以用 EVT(Epilogue Visitor Tree) 把 epilogue 写成一棵可组合的算子树。
13.6 把 GEMM 拼起来
完整的 CUTLASS 3.x GEMM 调用:
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
Shape<int, int, int, int>, // ProblemShape (M, N, K, L), L 是 batch 数
CollectiveMainloop, // 上面定义的 mainloop
CollectiveEpilogue // 上面定义的 epilogue
>;
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
// stride_A 等由 cutlass::make_cute_packed_stride 构造(cutlass-4.7.0/tools/util/include/cutlass/util/packed_stride.hpp),
// workspace 大小由 Gemm::get_workspace_size(args) 给出;这里从略
Gemm gemm;
typename Gemm::Arguments args = {
cutlass::gemm::GemmUniversalMode::kGemm,
{M, N, K, 1},
{ptr_A, stride_A, ptr_B, stride_B},
{{alpha, beta}, ptr_C, stride_C, ptr_D, stride_D}
};
gemm.initialize(args, workspace, stream);
gemm.run(stream);
这就是用户视角的 CUTLASS 3.x。看起来仍然有很多模板,但每个模板参数都有清晰含义。
13.7 怎么读 CUTLASS 源码
最后给读者一份"CUTLASS 源码导航指南"。GitHub 仓库 https://github.com/NVIDIA/cutlass,关键目录:
include/cutlass/gemm/:GEMM 主体。device/:device-level API(用户调用层)。kernel/:kernel-level(GemmUniversal 这一层)。collective/:CollectiveOp 实现。threadblock//warp/:2.x 那一套的底层 building block。
include/cutlass/epilogue/:epilogue 实现。include/cute/:CuTe 子库。layout.hpp:Layout 定义。tensor.hpp:Tensor 抽象。algorithm/copy.hpp:copy 算法(默认按 layout 自动向量化;cp.async / TMA 通过 Copy Atom 指定)。
examples/:可工作的示例代码。最简单的是examples/00_basic_gemm(2.x API)。
阅读建议:
- 从 example 开始:
cutlass-4.7.0/examples/48_hopper_warp_specialized_gemm是 Hopper 上的标准例子。读它的 main + 内嵌的 collective 配置。 - 追到 mainloop:从 example 进入
cutlass::gemm::collective::CollectiveBuilder,看它如何根据模板参数选择具体的 collective 实现(SM90 + WarpSpecialized路径)。 - 看 mainloop 体:找到
cutlass-4.7.0/include/cutlass/gemm/collective/sm90_mma_tma_gmma_ss_warpspecialized.hpp(ss= A、B 都从 SMEM 喂给 WGMMA;另有_rs_变体表示 A 走寄存器;第 9 章提到的 Machete 就是从其中的sm90_mma_tma_gmma_rs_warpspecialized_mixed_input.hpp改出来的,见vllm-0.8.5/csrc/quantization/machete/machete_mainloop.cuh:3),读它的load()(:310)/mma()(:417)两个成员——producer 与 consumer 的主体分别在这两处。 - 看 CuTe 操作:mainloop 里大量的
copy、gemm、partition_*都是 CuTe 函数,跳到include/cute看。
第一遍读会很慢(一个 example 一周)。但读懂之后回头看 FA3 源码、vLLM 的 Machete,会发现它们都直接建在 CUTLASS / CuTe 上,风格一致。Marlin 则是另一条路:vllm-0.8.5/csrc/quantization/gptq_marlin/ 下不引任何 CUTLASS 头文件,mma 直接写内联 PTX,拿来对照正好能看出 CUTLASS 替你省掉了什么。cuBLAS 闭源,看不到实现,只能用 Nsight 看它实际发了哪些 kernel。
13.8 第三篇收官:从 GFLOPs 到 TFLOPs
第三篇我们走完了 GEMM 的优化全程。本专栏没有条件在 H100 上逐版实测,下表沿用第 10–12 章的口径,只列各章引用过的公开实测和官方峰值:
| 章节 | 路径 | 数据类型 | 公开参照 / 天花板 |
|---|---|---|---|
| 第 10 章 | 朴素 GEMM | FP32 SIMT | 约为 cuBLAS SGEMM 的 8.5%(Boehm 博文,RTX A6000) |
| 第 11 章 | Tiled GEMM(SIMT 极限) | FP32 SIMT | 手写可逼近 cuBLAS SGEMM(Boehm 的 warptiling 版 93.7%,RTX A6000);通路峰值 67 TFLOPS(H100 SXM5) |
| 第 12 章 | mma.sync HGEMM 骨架 | FP16 TC | mma.sync 指令本身实测只到峰值约 64.9%(Luo et al. 2024 Table VII,H800) |
| 第 13 章 | CUTLASS 3.x(TMA + WGMMA) | FP16 TC | 稠密峰值 989 TFLOPS(H100 SXM5);本专栏没有同口径的 GEMM 实测可引 |
注意这张表跨了两条不同的赛道:前两行是 FP32 SIMT(H100 SXM5 峰值 67 TFLOPS),后两行是 FP16 Tensor Core(稠密峰值 989 TFLOPS),两条赛道的天花板差约 15 倍。同一条赛道里,差距来自 tiling、访存和流水工艺——朴素写法离 cuBLAS SGEMM 差一个数量级以上,手写 SIMT 却能逼近它;换到 Tensor Core 是换赛道;而在 Hopper 上想越过 mma.sync 那六成多的上限,就得换成 TMA + WGMMA,再配上 CUTLASS 级别的流水与调度,也就是本章讲的这一套。这就是"会写 CUDA"和"懂现代 GPU"的分野。
但 GEMM 不是最终目的。LLM 推理真正的难点是 Attention——它内部包含两个 GEMM(QK^T、PV),中间夹一个 softmax,且数据形态特殊(causal mask、长序列、KV cache)。
第 14-18 章的第四篇我们会把整个第二、三篇的工艺集中到一件事上:手写 FlashAttention v2 到 SOTA。从访存瓶颈分析(第 14 章),到 FA2 前向(第 15 章),到反向(第 16 章),到 Hopper 上的 TMA + Warp Specialization 优化(第 17 章),到 Persistent Kernel(第 18 章)。读完第 18 章,读者会拥有完整的 FA2 实现能力,并能看懂 FlashAttention v3 论文的主要创新。
本章动手练习:
- 在 H100 上跑 CUTLASS 的
cutlass-4.7.0/examples/48_hopper_warp_specialized_gemm。用--m=4096 --n=11008 --k=4096把问题尺寸换成 LLaMA-7B 的 MLP GEMM(4096 个 token),看实际性能。- 阅读 CuTe 的
Layout定义,理解Shape和Stride的嵌套语法。试着用 CuTe 描述一个 16×16 fp16 矩阵的 swizzle 布局。- 找到 CUTLASS 中 LinearCombination Epilogue 的实现,理解它是如何在最后阶段做 alpha/beta 的。