CUDA 算子工程:手写 FlashAttention v2 之路
第 10 章 朴素 GEMM 与 Roofline 分析
朴素 GEMM 的问题从来不是"算力不够"或"带宽不够"。 它的理论算术强度高达 683 FLOPs/byte,远在 H100 FP32 通路的临界点 20 右边—— 跑不快,是因为复用全交给了 cache:每做一次乘加就要发两条全局 load,按访存指令算,算术强度只有 0.25。
10.1 GEMM 是什么,为什么它占据 LLM 算力的 90%
GEMM = General Matrix Multiplication,通用矩阵乘:
LLM 推理 / 训练中无处不在:
- QKV projection:
X @ W_qkv(B×H × H×3H) - Attention 中的 Q@K^T 与 P@V:每个 head 一次 (Seq_q×d × d×Seq_k)、(Seq_q×Seq_k × Seq_k×d)
- Output projection:
O @ W_o - FFN gate / up / down:3 个独立 GEMM
一个 LLaMA-7B 单 token decoding 的浮点运算分布,可以按模型形状粗估(d=4096、32 层、FFN 中间维 11008、词表 32000,上下文长度 2048;一次乘加记 2 FLOPs):
权重 GEMM (QKV/O/FFN/LM head): ~92% (约 13.2 GFLOPs/token)
Attention 的 QK^T 与 PV: ~7.5% (随上下文线性增长:512 时约 2%,4096 时约 14%)
RMSNorm、SiLU、softmax 等: <0.2%
attention 的 QK^T 与 PV 本身也是矩阵乘,把它们也算进去,99% 以上的浮点运算都在矩阵乘上。所以 GEMM 性能直接决定 LLM 的算力利用率。(注意这是浮点运算的占比,不是耗时的占比:decode 阶段的 GEMM 往往被读权重的带宽卡住,归一化、激活这些小算子也要各读写一遍显存。)
大尺寸下(比如 M=N=K=4096),cuBLAS 的 SGEMM 和 Tensor Core HGEMM 都能吃掉各自峰值的大半——H100 SXM5 官方规格里,FP32 非 Tensor Core 峰值是 67 TFLOPs,FP16/BF16 Tensor Core 稠密峰值是 989 TFLOPs。这个性能不是凭空来的——它是多代 GPU 架构演进加上大量 kernel 工程堆出来的,其中许多技巧已经由 CUTLASS 以开源模板的形式公开。
要理解 GEMM 性能的来路,得从最简单的版本开始。
10.2 朴素 GEMM:教科书式的三重循环
最直观的写法:每个线程算 C 的一个元素。
__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) return;
float sum = 0.0f;
for (int k = 0; k < K; ++k) {
sum += A[row * K + k] * B[k * N + col];
}
C[row * N + col] = sum;
}
启动配置:
dim3 block(16, 16);
dim3 grid((N + 15) / 16, (M + 15) / 16);
gemm_naive<<<grid, block>>>(A, B, C, M, N, K);
这段代码(对应 α=1、β=0)逻辑正确,能算对 GEMM。但性能呢?
10.2.1 它慢在哪个量级
手头没有 H100 上同口径的公开对照数,这里借 Simon Boehm 的博文 How to Optimize a CUDA Matmul Kernel for cuBLAS-like Performance: a Worklog(2022 年 12 月)作参照。注意他用的 GPU 是 RTX A6000(Ampere),不是 H100,矩阵是 4092×4092 的 FP32:
朴素,threadIdx.x 对应 C 的行: 309 GFLOPs/s (cuBLAS 的 1.3%)
朴素,threadIdx.x 对应 C 的列: 1986 GFLOPs/s (cuBLAS 的 8.5%)
cuBLAS SGEMM: 23250 GFLOPs/s
本章的 gemm_naive 让 threadIdx.x 对应列,属于第二行的写法。
这里刻意不给 H100 的绝对数字:本专栏写作时没有条件在 H100 上重跑这组对比,而 GEMM 的绝对数字对 GPU 型号、CUDA 版本、cuBLAS 版本、时钟策略都很敏感。可以放心带走的是量级结论——朴素写法离 cuBLAS 差一个数量级以上,线程映射写反了还要再差几倍。原因下面用 Roofline 和访存模式来分析。
一个数量级以上的差距。这就是"会写"和"写好"的距离。
10.3 用 Roofline 找瓶颈
朴素 GEMM 的瓶颈在哪?用第 4 章学的 Roofline 模型分析。
10.3.1 朴素 GEMM 的算术强度
每个线程计算 C 的一个元素,需要:
- 读 A 的一行: 个 float = 字节
- 读 B 的一列: 个 float = 字节
- 写 C 的一个元素:4 字节
- 计算: 个 FLOPs(K 次 mul + K 次 add)
每个 thread 的算术强度:
K 大时趋近于 1/4。这是非常低的算术强度——和 LayerNorm、RMSNorm 这类纯访存算子一个量级,甚至更低。
但整个 kernel 的算术强度不是这样算的。多个线程之间会复用 A 和 B 的不同部分:
- 同一行线程都读同一行 A
- 同一列线程都读同一列 B
如果 cache 完美命中,整个 kernel 的总访存只需要:
- A 全部读 1 次: 字节
- B 全部读 1 次: 字节
- C 全部写 1 次: 字节
- 计算: FLOPs
整体算术强度:
M=N=K=4096 时:
683 FLOPs/byte 的算术强度,远超过 H100 SXM5 FP32 通路的临界点 20(67 TFLOPs ÷ 3.35 TB/s),甚至超过第 4 章那个 FP16/BF16 Tensor Core 稠密的临界点 295(989 ÷ 3.35)。这意味着 GEMM 在大尺寸下应该是完全 compute-bound 的——理论上能跑到算力峰值。
但朴素 GEMM 离峰值差得很远。问题不是算力或带宽不够,而是数据复用全交给了 cache,程序自己一点没安排。
10.3.2 朴素 GEMM 为什么复用不起来
看朴素 GEMM 的访存模式:
sum += A[row * K + k] * B[k * N + col];
先按 warp 看合并(口径同第 1、4 章:按 32 字节 sector 计)。block(16, 16) 下 threadIdx.x 对应列,一个 warp 的 32 个线程是 2 行 × 16 列:
- A:
A[row * K + k]与 col 无关,同一行的 16 个线程读同一个地址,走广播;一条 load 指令只涉及 2 个地址、2 个 sector。 - B:
B[k * N + col]里 col 连续,每行 16 个线程读连续的 64 字节,两行线程读的还是同一段,一共 2 个 sector(B 按 64 字节对齐时;不对齐也只多 1 个)。
所以这一版的合并本身没有问题。反过来,如果让 threadIdx.x 对应行(Boehm 博文里最初的朴素版就是这样),warp 内读 A 的 32 个地址两两相距 K×4 字节,一条 load 散成 32 个 sector,就是上面 1.3% 那一档。
真正的问题在复用:
- 单个线程沿 k 往前走时,读 A 是连续的,读 B 每一步却跨 N×4 字节(N=4096 时 16 KB,即 128 个 128 字节的 cache line)。线程自己的访问谈不上局部性,只能指望别的线程已经把同一段数据带进了 cache。
- 每个 A 元素要被 N 个线程各读一次,每个 B 元素要被 M 个线程各读一次。load 与 FFMA 的比例是 2:1——nvcc 13.4 为 sm_90a 编译出的 SASS 里,主循环正是每条 FFMA 配两条 LDG(不同 CUDA 版本可能不同)。
- 这些重复读取能被吸收多少,全看 L1/L2:同一个 block 的 8 个 warp 读的是同一段 B,有机会在 L1 命中;block 之间只能指望 L2——H100 SXM5 的 L2 是 50 MB,而 N=4096 时单是 B 就有 64 MiB,放不下整块,能否命中取决于同时在跑的那批 block 是否恰好在读同一片数据。
结果是:按访存指令口径,算术强度只有 0.25 FLOPs/byte;即使全部命中 cache,每条 FFMA 也要陪两条 LDG 走一遍 LSU 和 L1,访存通路会先于 FFMA 单元成为瓶颈。最终落到 HBM 上的流量有多少、按 HBM 口径的算术强度是多少,取决于各级 cache 的命中率,只能用 Nsight Compute 实测(见本章练习 2)。
10.3.3 朴素 GEMM 的 Roofline 位置
实际 TFLOPs/s
▲
989 ─────┼────────────────────── FP16/BF16 Tensor Core 稠密峰值
│
100 ─┤
67 ─────┼────────────────────── FP32 非 Tensor Core 峰值
│ cuBLAS SGEMM (FP32)
│ AI~683, 贴近 FP32 算力墙
10 ─┤ ╱
│ ╱
│ 朴素 GEMM:复用全靠 cache,
│ 离 FP32 算力墙差一个数量级以上
1 ──┤
└─────────────────────► 算术强度 (对 HBM)
1 10 100 1000
斜线:HBM 带宽屋顶 (H100 SXM5 HBM3 3.35 TB/s),与 67 的交点在 AI=20
以上峰值均为 H100 SXM5 官方规格。朴素 GEMM 把复用全交给 cache,按访存指令算的算术强度只有 0.25,所以无法发挥算力优势。优化的方向是把复用从 cache 的偶然命中变成程序显式的安排——把 A/B 的 tile 搬进 SMEM 和寄存器反复使用,让实际打到各级存储上的算术强度逼近理论的 683。
10.4 GEMM 优化的三大武器
要让 GEMM 跑到算力峰值,需要三层 tile:
flowchart TB HBM[A, B 在 HBM] -->|每次 K_tile 列| GTile[Block-level tile] GTile -->|放进 SMEM| SMEM[SMEM tile A, B] SMEM -->|每次 mma_k 列| WTile[Warp-level tile] WTile -->|放进寄存器| REG[Register fragment] REG -->|喂给 Tensor Core| TC[mma.sync 一次<br/>16×8×16 矩阵乘] TC -->|累加| ACC[寄存器中的 C accumulator]
三层 tile:
- Block-level tile:一个 Block 处理 C 的一个 M_tile × N_tile 子块(典型 128×128)。一个 SM 上能同时驻留几个 block,取 SMEM、寄存器、线程数等各项上限里最紧的那个;SMEM 最紧时约为 SMEM_total / block_smem_size。
- Warp-level tile:一个 warp 处理 block tile 内的一个 M_warp × N_warp 子块(第 11 章用 64×32,128×128 的 block tile 正好分给 8 个 warp)。
- Tensor Core mma:一个 mma.sync 操作 16×8×16(Hopper 上的 WGMMA 是 64×N×16)。
这三层 tile 配合 double buffering(异步拷贝下一个 K tile 时算当前),构成现代 GEMM 的核心结构。
10.4.1 数据复用计算
理想情况下:
- Block tile (128×128):每个 K 列的 A_tile (128) + B_tile (128) 共 256 个 float = 1 KB,被用来算 128×128 = 16384 个 C 元素。复用率 16384 / 256 = 64×。
- Warp tile (64×32):每个 K 列从 SMEM 读 64 + 32 = 96 个元素,做 64×32 = 2048 次乘加,复用率约 21×。
- Register tile (8×8):每个 K 列读进寄存器的 8 + 8 个元素做 64 次乘加,每个元素被用 8 次。
这种层层复用让每个从 HBM 读上来的字节被多次使用,有效算术强度 接近理论上限。
10.4.2 Tensor Core 的角色
第 12 章会详细讲,这里先建立直觉:
- Ampere 起,一条 mma.sync.m16n8k16 指令算 16×8×16 矩阵乘(Volta 的 mma 形状是 m8n8k4)。
- 算 16×8×16 矩阵乘需要的浮点操作 = 16 × 8 × 16 × 2 = 4096 FLOPs。
- 一个 SM 有 4 个 Tensor Core,但一条 m16n8k16 并不是一个周期就吐出结果——按官方峰值反推,四代 Tensor Core 上它占用约 4 个周期。
- 折算下来每 SM 每周期约 4096 FLOPs(4 个 TC × 4096 FLOPs ÷ 4 周期)。作为对照,A100(三代 TC)是 2048,Hopper 正好翻倍。
- 132 SM × 1.83 GHz × 4096 FLOPs/周期 ≈ 989 TFLOPs/s。
这就是 H100 SXM5 FP16/BF16 Tensor Core 稠密峰值的来路(1.83 GHz 是由官方峰值反推的时钟,见第 1 章的说明)。要跑满这个峰值,每个 SM 每周期必须发出一条 mma 指令——这要求所有数据准备好、SMEM/寄存器完美配合。任何一个环节 stall(cache miss、bank conflict、寄存器 spill),峰值就打不到。
10.5 优化路径预告
后三章我们会一步步从朴素 GEMM 推进到 cuBLAS 性能:
- 第 11 章 Tiled GEMM:用 SMEM tile + double buffer,实现 SIMT GEMM。不用 Tensor Core 的手写版本可以逼近 cuBLAS SGEMM——Boehm 在 RTX A6000 上的实测里,warptiling 版达到 cuBLAS 的 93.7%。
- 第 12 章 Tensor Core GEMM:用 mma.sync + ldmatrix + swizzle,引入 Tensor Core,算力天花板从 FP32 通路的 67 TFLOPS 升到 FP16 稠密 989 TFLOPS(H100 SXM5 官方规格)。
- 第 13 章 CUTLASS 设计哲学:剖析 NVIDIA 官方模板库的设计,理解工业级 GEMM 是怎么组装的。
每一步都对应一组真实的工程技巧,每一组技巧都把性能往上推一截。这条优化路径同时也是后续 FA2 优化的基础——FA2 内部的 QK^T 和 PV 矩阵乘,都基于这套 GEMM 框架。
10.6 一个朴素优化:Block 内复用
在进入第 11 章之前,先做一个小优化作为"warm-up"——让 thread 算多个 C 元素:
// 启动:dim3 block(16, 16);
// dim3 grid((N + 16 * TN - 1) / (16 * TN), (M + 16 * TM - 1) / (16 * TM));
template <int TM = 4, int TN = 4> // 每线程算 TM × TN 个 C 元素
__global__ void gemm_thread_tile(
const float* A, const float* B, float* C,
int M, int N, int K
) {
int row = (blockIdx.y * blockDim.y + threadIdx.y) * TM;
int col = (blockIdx.x * blockDim.x + threadIdx.x) * TN;
float c[TM][TN] = {0};
for (int k = 0; k < K; ++k) {
float a[TM], b[TN];
#pragma unroll
for (int i = 0; i < TM; ++i) a[i] = (row + i < M) ? A[(row + i) * K + k] : 0.0f;
#pragma unroll
for (int j = 0; j < TN; ++j) b[j] = (col + j < N) ? B[k * N + (col + j)] : 0.0f;
#pragma unroll
for (int i = 0; i < TM; ++i)
#pragma unroll
for (int j = 0; j < TN; ++j)
c[i][j] += a[i] * b[j];
}
#pragma unroll
for (int i = 0; i < TM; ++i)
#pragma unroll
for (int j = 0; j < TN; ++j)
if (row + i < M && col + j < N) // M、N 不是 16×TM、16×TN 的整数倍时防越界
C[(row + i) * N + (col + j)] = c[i][j];
}
每线程算 4×4=16 个 C 元素。每个 k 读 A 的 4 个元素 + B 的 4 个元素 = 8 次访存,做 16 次 mul-add。按访存指令口径,算术强度变成 32 FLOPs / 32 字节 = 1 FLOPs/byte——比朴素的 0.25 提升 4 倍。nvcc 13.4 为 sm_90a 编译出的 SASS 里,每个 k 正是 8 条 LDG 配 16 条 FFMA,LDG 与 FFMA 之比从 2:1 变成 1:2(不同 CUDA 版本可能不同)。
访存指令少了四分之三,对被访存通路卡住的 kernel 来说这是实打实的收益——这正是 Roofline 的用法:受访存限制时,性能随(对应那一层存储的)算术强度成正比上升。但涨幅要实测,因为 cache 行为也变了:col = ... * TN 让 warp 内相邻线程读 B 的地址相距 16 字节,一条 LDG 涉及的 sector 数是朴素版的 4 倍,要靠紧随其后 j=1..3 的几条 LDG 在 L1 命中找补回来。
但这远没到目标。下一步要做的是让多个 thread 共享 A/B——这就是 Block tile + SMEM 的舞台。
10.7 这一章的小结与下一章
这一章我们建立了 GEMM 优化的"地图":
- 矩阵乘占 LLM 浮点运算的 99% 以上(含 attention 内部的 QK^T、PV):所有优化的最大投资点。
- 朴素 GEMM 离 cuBLAS 差一个数量级以上:不是算力或带宽不够,而是复用全交给了 cache。
- 理论算术强度高(~683),按访存指令算只有 0.25:必须靠 SMEM + 寄存器层层复用提升实际 AI。
- 三层 tile 是 GEMM 优化的核心结构:Block tile / Warp tile / Tensor Core mma。
- 简单的 thread tile 就把访存指令口径的算术强度从 0.25 提到 1:访存指令少了四分之三,但还远没到峰值,需要 SMEM 和 Tensor Core。
第 11 章我们正式进入 SMEM 优化的世界——把 Block tile 放进 SMEM,配合 double buffering 让 SMEM 和 HBM 流水起来。这个版本(不用 Tensor Core)能把 FP32 通路用到相当充分的程度,是 SIMT GEMM 的合理目标。第 12 章换上 Tensor Core 后,天花板从 FP32 通路的 67 TFLOPS 抬到 FP16 稠密 989 TFLOPS,手写 HGEMM 骨架就能把 SIMT 写法甩开一个数量级;不过 Hopper 上 mma.sync 本身只能跑到峰值的约 65%(Luo et al. 2024 在 H800 上实测),再往上要靠 WGMMA。
本章动手练习:
- 实现朴素 GEMM 和 thread tile 4×4 GEMM,对比性能。
- 用 Nsight Compute 看朴素 GEMM 的
dram__sectors_read.sum(每个 sector 32 字节),估算实际 HBM 流量,算出按 HBM 口径的算术强度,与按访存指令口径的 0.25 对比,看 L1/L2 吸收了多少重复读取。- 思考:为什么 thread tile 要选 4×4 而不是 8×8?(提示:寄存器压力。nvcc 13.4 为 sm_90a 编译上面带边界检查的版本,4×4 用 39 个寄存器,8×8 用 114 个,不同 CUDA 版本可能不同)