CUDA 算子工程:手写 FlashAttention v2 之路

第 10 章 朴素 GEMM 与 Roofline 分析

作者 杨艺韬 · 3,707 字 · 发布于 · 更新于

朴素 GEMM 的问题从来不是"算力不够"或"带宽不够"。 它的理论算术强度高达 683 FLOPs/byte,远在 H100 FP32 通路的临界点 20 右边—— 跑不快,是因为复用全交给了 cache:每做一次乘加就要发两条全局 load,按访存指令算,算术强度只有 0.25。

10.1 GEMM 是什么,为什么它占据 LLM 算力的 90%

GEMM = General Matrix Multiplication,通用矩阵乘:

Cm,n=α∑k=1KAm,kBk,n+βCm,nC_{m,n} = \alpha \sum_{k=1}^{K} A_{m,k} B_{k,n} + \beta C_{m,n}

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 的一行:KK 个 float = 4K4K 字节
  • 读 B 的一列:KK 个 float = 4K4K 字节
  • 写 C 的一个元素:4 字节
  • 计算:2K2K 个 FLOPs(K 次 mul + K 次 add)

每个 thread 的算术强度:

AIthread=2K4K+4K+4=2K8K+4≈14 FLOPs/byte\text{AI}_{\text{thread}} = \frac{2K}{4K + 4K + 4} = \frac{2K}{8K + 4} \approx \frac{1}{4} \text{ FLOPs/byte}

K 大时趋近于 1/4。这是非常低的算术强度——和 LayerNorm、RMSNorm 这类纯访存算子一个量级,甚至更低。

但整个 kernel 的算术强度不是这样算的。多个线程之间会复用 A 和 B 的不同部分:

  • 同一行线程都读同一行 A
  • 同一列线程都读同一列 B

如果 cache 完美命中,整个 kernel 的总访存只需要:

  • A 全部读 1 次:M×K×4M \times K \times 4 字节
  • B 全部读 1 次:K×N×4K \times N \times 4 字节
  • C 全部写 1 次:M×N×4M \times N \times 4 字节
  • 计算:2×M×N×K2 \times M \times N \times K FLOPs

整体算术强度:

AIkernel=2MNK4(MK+KN+MN)\text{AI}_{\text{kernel}} = \frac{2 M N K}{4(MK + KN + MN)}

M=N=K=4096 时:

AI=2⋅409634⋅3⋅40962=2⋅409612≈683 FLOPs/byte\text{AI} = \frac{2 \cdot 4096^3}{4 \cdot 3 \cdot 4096^2} = \frac{2 \cdot 4096}{12} \approx 683 \text{ FLOPs/byte}

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:

  1. Block-level tile:一个 Block 处理 C 的一个 M_tile × N_tile 子块(典型 128×128)。一个 SM 上能同时驻留几个 block,取 SMEM、寄存器、线程数等各项上限里最紧的那个;SMEM 最紧时约为 SMEM_total / block_smem_size。
  2. Warp-level tile:一个 warp 处理 block tile 内的一个 M_warp × N_warp 子块(第 11 章用 64×32,128×128 的 block tile 正好分给 8 个 warp)。
  3. 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 优化的"地图":

  1. 矩阵乘占 LLM 浮点运算的 99% 以上(含 attention 内部的 QK^T、PV):所有优化的最大投资点。
  2. 朴素 GEMM 离 cuBLAS 差一个数量级以上:不是算力或带宽不够,而是复用全交给了 cache。
  3. 理论算术强度高(~683),按访存指令算只有 0.25:必须靠 SMEM + 寄存器层层复用提升实际 AI。
  4. 三层 tile 是 GEMM 优化的核心结构:Block tile / Warp tile / Tensor Core mma。
  5. 简单的 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。

本章动手练习:

  1. 实现朴素 GEMM 和 thread tile 4×4 GEMM,对比性能。
  2. 用 Nsight Compute 看朴素 GEMM 的 dram__sectors_read.sum(每个 sector 32 字节),估算实际 HBM 流量,算出按 HBM 口径的算术强度,与按访存指令口径的 0.25 对比,看 L1/L2 吸收了多少重复读取。
  3. 思考:为什么 thread tile 要选 4×4 而不是 8×8?(提示:寄存器压力。nvcc 13.4 为 sm_90a 编译上面带边界检查的版本,4×4 用 39 个寄存器,8×8 用 114 个,不同 CUDA 版本可能不同)