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

第 11 章 Tiled GEMM:Shared Memory 与 Double Buffer

作者 杨艺韬 · 4,324 字 · 发布于 · 更新于

朴素 GEMM 把复用全交给 cache,按访存指令算,算术强度只有 0.25,而 H100 SXM5 FP32 通路的 Roofline 临界点约 20(67 TFLOPs ÷ 3.35 TB/s)。 这一章做的全部事情,就是用 block / warp / thread 三层 tile 把复用变成程序显式的安排,把这个数字推过临界点。

11.1 为什么需要 Tiling

第 10 章看到朴素 GEMM 离 cuBLAS 差一个数量级以上,症结是复用全交给了 cache:每做一次乘加就要发两条全局 load,按访存指令算,算术强度只有 0.25。要让算力发挥出来,必须显式地把 A 和 B 的元素复用起来——同一个元素从 HBM 读进 SMEM 后,被多个线程多次使用。

Tiling(瓦片化)就是这种复用的工程实现。基本结构是:

flowchart TB
  subgraph HBM_Layer [HBM 层]
    A[A 矩阵 M×K]
    B[B 矩阵 K×N]
    C[C 矩阵 M×N]
  end
  subgraph SMEM_Layer [SMEM 层 · per Block]
    SA[A_tile 128×Tk]
    SB[B_tile Tk×128]
  end
  subgraph Reg_Layer [Register 层 · per Warp]
    RA[A_frag 64×Tk_inner]
    RB[B_frag Tk_inner×32]
    RC[C_acc 64×32]
  end
  HBM_Layer -->|Block iteration<br/>每次 K 维移动 Tk| SMEM_Layer
  SMEM_Layer -->|Warp iteration<br/>每次 K 维移动 Tk_inner| Reg_Layer
  Reg_Layer -->|FMA 累加| Reg_Layer

每个 Block 处理 C 的一个 M_block × N_block 子块(典型 128×128),扫过 K 维度时不断从 HBM 加载新的 A_tile / B_tile 到 SMEM。每个 Warp 处理 block 内的一个 M_warp × N_warp 子块,从 SMEM 读 fragment 到寄存器。最内层是寄存器中的 8×8 累加(每线程)。这三层尺寸不能随便凑——下一节会看到它们被一条等式锁死。

11.2 Block Tile:SMEM 中的复用

定义 tile 大小:

constexpr int BM = 128;  // Block tile M 维
constexpr int BN = 128;  // Block tile N 维
constexpr int BK = 16;   // Block tile K 维

每个 Block 处理 BM × BN 的 C 子块。内层 K 维分块 BK,每次从 HBM 拉 BM × BK 的 A_tile 和 BK × BN 的 B_tile 到 SMEM。

SMEM 占用:

A_tile: 128 × 16 × 4 = 8192 bytes = 8 KB
B_tile: 16 × 128 × 4 = 8192 bytes = 8 KB
合计:                   16 KB

H100 SXM5 每 SM 最多 228 KB SMEM,每个 block 另有 1 KB 系统预留,单看 SMEM 能驻留 13 个 block(228 ÷ 17)。但真正卡住的是别处:每 SM 最多 2048 个线程,256 线程的 block 最多 8 个;寄存器更紧——nvcc 13.4 为 sm_90a 编译 11.3 节的 kernel,每线程 96 个寄存器,一个 block 要 24576 个,每 SM 65536 个寄存器只够驻留 2 个 block(不同 CUDA 版本可能不同)。

11.2.1 加载 A_tile / B_tile 到 SMEM

每个 Block 启动 BM × BN / (TM × TN) = 256 个线程(每线程算 8×8 = 64 个 C 元素)。256 线程协作加载 8 KB 的 A_tile 和 8 KB 的 B_tile。

__shared__ float sA[BM][BK];   // 128 × 16
__shared__ float sB[BK][BN];   // 16 × 128

const int tid = threadIdx.y * blockDim.x + threadIdx.x;
const int row = blockIdx.y * BM;
const int col = blockIdx.x * BN;

// A_tile 加载: 256 线程加载 128 × 16 = 2048 个元素
// 每线程加载 2048 / 256 = 8 个元素 (但 128 行 × 16 列, 安排成每 16 线程一行)
for (int load = 0; load < BM * BK / 256; load++) {
    int idx = tid + load * 256;
    int load_row = idx / BK;
    int load_col = idx % BK;
    sA[load_row][load_col] = A[(row + load_row) * K + (k_step + load_col)];
}

类似地加载 B_tile。

11.2.2 内层乘法:从 SMEM 到寄存器

三层尺寸必须自洽:每线程算 TM × TN 个 C 元素,一个 warp 32 线程,所以 warp tile = 32 × TM × TN;一个 block 8 个 warp(256 线程),所以 block tile = 8 × warp tile。代进 TM=TN=8、BM=BN=128:

warp tile 面积 = 32 × 64 = 2048 = 64 × 32     → WM=64, WN=32
block tile 面积 = 8 × 2048 = 16384 = 128 × 128 ✓
warp 网格 = (BM/WM) × (BN/WN) = 2 × 4 = 8 warp ✓
warp 内线程网格 = (WM/TM) × (WN/TN) = 8 × 4 = 32 lane ✓

注意 warp tile 是 64×32 而不是 64×64——后者需要 4096 个元素、也就是每线程 128 个累加器,一个 block 只剩 4 个 warp,与 256 线程的划分对不上。这类"面积对不上"是手写 GEMM 最常见的低级错误。

constexpr int WM = 64;   // Warp tile M
constexpr int WN = 32;   // Warp tile N
constexpr int TM = 8;    // Thread tile M
constexpr int TN = 8;    // Thread tile N

float c[TM][TN] = {0};   // 64 个寄存器

for (int kk = 0; kk < BK; ++kk) {
    float a[TM], b[TN];
    #pragma unroll
    for (int i = 0; i < TM; ++i) a[i] = sA[warp_m * WM + thread_m * TM + i][kk];
    #pragma unroll
    for (int j = 0; j < TN; ++j) b[j] = sB[kk][warp_n * WN + thread_n * TN + j];

    #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];
}

每个线程从 SMEM 读 8 个 a + 8 个 b = 16 个浮点(64 字节),做 64 次 mul-add(128 FLOPs)。单线程算术强度 = 128 / 64 = 2 FLOPs/byte(按访存指令口径,这里访问的是 SMEM)——是第 10 章朴素版 0.25 的 8 倍,4×4 thread tile 的 2 倍。

11.3 完整的 Tiled GEMM Kernel

把上面的 piece 组装起来:

template <int BM, int BN, int BK, int WM, int WN, int TM, int TN>
__global__ void gemm_tiled(
    const float* A, const float* B, float* C,
    int M, int N, int K
) {
    __shared__ float sA[BM][BK];
    __shared__ float sB[BK][BN];

    const int tid = threadIdx.y * blockDim.x + threadIdx.x;
    const int warp_id = tid / 32;
    const int lane_id = tid % 32;
    const int warp_m = warp_id / (BN / WN);   // BN/WN = 4, 所以 warp_m ∈ [0,2)
    const int warp_n = warp_id % (BN / WN);   //            warp_n ∈ [0,4)
    const int thread_m = lane_id / (WN / TN);
    const int thread_n = lane_id % (WN / TN);

    const int block_row = blockIdx.y * BM;
    const int block_col = blockIdx.x * BN;

    float c[TM][TN] = {0};

    // 沿 K 维度迭代
    for (int k_step = 0; k_step < K; k_step += BK) {
        // 1. 协作加载 A_tile, B_tile 到 SMEM
        #pragma unroll
        for (int i = tid; i < BM * BK; i += blockDim.x * blockDim.y) {
            int r = i / BK, c_ = i % BK;
            int gr = block_row + r, gc = k_step + c_;
            sA[r][c_] = (gr < M && gc < K) ? A[gr * K + gc] : 0.0f;   // 越界补 0
        }
        #pragma unroll
        for (int i = tid; i < BK * BN; i += blockDim.x * blockDim.y) {
            int r = i / BN, c_ = i % BN;
            int gr = k_step + r, gc = block_col + c_;
            sB[r][c_] = (gr < K && gc < N) ? B[gr * N + gc] : 0.0f;
        }
        __syncthreads();

        // 2. 每线程算 TM × TN 个累加
        for (int kk = 0; kk < BK; ++kk) {
            float a[TM], b[TN];
            #pragma unroll
            for (int i = 0; i < TM; ++i)
                a[i] = sA[warp_m * WM + thread_m * TM + i][kk];
            #pragma unroll
            for (int j = 0; j < TN; ++j)
                b[j] = sB[kk][warp_n * WN + thread_n * TN + j];

            #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];
        }
        __syncthreads();
    }

    // 3. 写 C
    #pragma unroll
    for (int i = 0; i < TM; ++i) {
        #pragma unroll
        for (int j = 0; j < TN; ++j) {
            int r = block_row + warp_m * WM + thread_m * TM + i;
            int c_ = block_col + warp_n * WN + thread_n * TN + j;
            if (r < M && c_ < N) C[r * N + c_] = c[i][j];
        }
    }
}

// Launch
dim3 block(16, 16);  // 256 线程
dim3 grid((N + BN - 1) / BN, (M + BM - 1) / BM);   // 向上取整
gemm_tiled<128, 128, 16, 64, 32, 8, 8><<<grid, block>>>(A, B, C, M, N, K);

加载时越界的 A/B 元素补 0、grid 向上取整,M、N、K 就不必是 tile 的整数倍;越界补的 0 对累加没有贡献,写回时再由 r < M && c_ < N 挡掉多算的部分。

本专栏没有条件在 H100 上实测。可作参照的公开实测来自 Simon Boehm 的博文 How to Optimize a CUDA Matmul Kernel for cuBLAS-like Performance: a Worklog(2022 年 12 月),GPU 是 RTX A6000(Ampere,NVIDIA 标称 FP32 峰值 38.7 TFLOPs),不是 H100,矩阵是 4092×4092 的 FP32:

朴素(threadIdx.x 对应 C 的列):        1986 GFLOPs/s    (cuBLAS 的 8.5%)
SMEM 缓存(每线程仍只算 1 个 C 元素):   2980 GFLOPs/s    (12.8%)
2D blocktiling(BM=BN=128, BK=8,
  TM=TN=8, 256 线程):                 15972 GFLOPs/s    (68.7%)
cuBLAS SGEMM:                         23250 GFLOPs/s

他的 2D blocktiling 与本节 kernel 的 tile 参数基本相同(只是 BK=8)。光把 tile 放进 SMEM 还不够——每线程只算一个元素时,SMEM 访问本身成了瓶颈(他的 profile 里主要停在 MIO Throttle,也就是 SMEM 指令队列满);加上 8×8 的 thread tile(中间还有一版每线程算一列 8 个元素的 1D 版,36.5%),才跨到 cuBLAS 的七成。

11.4 优化 1:Double Buffer + Async Copy

上面的代码有一个明显的"同步等待":每次 K 迭代完,要 __syncthreads() 等所有线程算完,才能加载下一个 K_tile。加载期间这个 block 的 warp 都在等数据,只能指望同一 SM 上的其他 block 来填空——而按 11.2 节的寄存器用量,一个 SM 只驻留 2 个 block。

Double buffer(双缓冲)让计算和加载重叠:

__shared__ float sA[2][BM][BK];   // 两份缓冲
__shared__ float sB[2][BK][BN];

// 设 K 是 BK 的整数倍(边界 tile 可用 cp.async 的 src-size 操作数补零,这里从略)
// 预加载第一个 buffer
load_to_smem(sA[0], sB[0], k_step=0);
__syncthreads();

for (int k_step = BK; k_step < K; k_step += BK) {
    int cur = ((k_step / BK) - 1) % 2;
    int next = (k_step / BK) % 2;

    // 异步加载下一个 buffer
    cp_async_global_to_shared(sA[next], A_addr_at(k_step));
    cp_async_global_to_shared(sB[next], B_addr_at(k_step));
    cp_async_commit();

    // 同时计算当前 buffer
    compute_on_smem(sA[cur], sB[cur], &c);

    cp_async_wait_all();
    __syncthreads();
}

// 算最后一个 buffer
compute_on_smem(sA[(K/BK - 1) % 2], sB[(K/BK - 1) % 2], &c);

这里用到了 Ampere(sm_80)起的 cp.async.cg.shared.global 指令——异步地从全局内存拷贝到 SMEM,期间 SIMT cores 可以继续算。它和第 3 章的手动双缓冲是同一个流水结构,差别在数据路径:第 3 章没有 cp.async,只能先把下一块读进寄存器、算完再写 SMEM(CUTLASS include/cutlass/gemm/threadblock/mma_pipelined.h 的寄存器中转做法);cp.async 由硬件直接把数据写进 SMEM,不经寄存器中转,所以发出拷贝后就能直接去算。把这段伪代码补全(每线程每次搬一个 16 字节的 float4)后用 nvcc 13.4 编译,sm_80 与 sm_90a 的结果一样:拷贝落成 LDGSTS.E.BYPASS.128,commit 与 wait_all 分别是 LDGDEPBAR 和 DEPBAR.LE SB0(不同 CUDA 版本可能不同)。.cg 即只缓存在 L2、绕过 L1(SASS 里的 BYPASS),拷贝大小必须是 16 字节。

这一步能省多少,取决于加载延迟原本被其他驻留 block 掩盖了多少。Boehm 的博文没有实现双缓冲,本专栏也没有可引用的同口径数字,这里不给具体比例。

11.5 优化 2:解决 SMEM Bank Conflict

读 sA[warp_m * WM + thread_m * TM + i][kk] 这一句。分析 bank conflict 有个前提必须先摆正:冲突是按"一条指令里 warp 的 32 个 lane"算的,不是按一个线程先后发的几条指令算的。这条 LDS 里 i、kk、warp_m 都是常量,随 lane 变的只有 thread_m = lane_id / 4,它取 0..7;thread_n 不进 A 的地址,所以每 4 个 lane 读的是同一个地址,那部分走广播、不算冲突。

于是一条指令里 warp 真正要取 8 个不同地址,行号是 base, base+8, ..., base+56——相邻两个差 TM = 8 行。每行 BK=16 个 float,跨 8 行就是 128 个 float,而 128 % 32 == 0:

lane  0.. 3 → sA[base+ 0][kk] → bank (16·base + kk) % 32
lane  4.. 7 → sA[base+ 8][kk] → bank 同上 (128 % 32 == 0)
lane  8..11 → sA[base+16][kk] → bank 同上
...
lane 28..31 → sA[base+56][kk] → bank 同上

8 个地址全撞在同一个 bank 上——标准的 8-way bank conflict,一条 LDS 要被拆成 8 次串行访问。

解决方法:+1 padding 或 swizzled 布局。

11.5.1 +1 Padding

__shared__ float sA[BM][BK + 1];   // 17 列而不是 16

行距从 16 变成 17 之后,跨 8 行是 136 个 float,136 % 32 == 8,8 个地址落到 +0 / +8 / +16 / +24 四个 bank 上、每个 bank 两个地址——8-way 降到 2-way。代价是每行多一列,浪费 SMEM 1/16 ≈ 6%。

这里要说破一件常被含糊过去的事:对这个访问模式,光靠行 padding 永远消不干净。warp 内相邻两个取数的行距恒为 TM=8,地址差是 8 × pitch 个 float;8 × pitch 必然是 8 的倍数,落进 32 个 bank 最多只能凑出 4 个不同的 bank,2-way 就是行 padding 的下限。要真正归零,动的必须是布局或线程映射——见下一节。

11.5.2 Swizzled Layout

更高级的做法是把列号按行号 XOR 一下:

__device__ __forceinline__ int swizzle(int row, int col) {
    // 关键:XOR 进去的必须是"warp 内真正在变的那几位行号"。
    // 这里 warp 内的 8 个行号相差 TM=8,变的是 bit 3..5,
    // 所以要先右移 3 位再取 3 位;写成 (row & 0x7) 是无效的——
    // base, base+8, ... base+56 的低 3 位完全相同,XOR 出来还是同一列。
    return col ^ ((row >> 3) & 0x7);
}
sA[i][swizzle(i, kk)] = ...;

这样 8 个地址的列号被打散成 8 个不同值,bank 也就跟着散开,冲突真正归零,而且一个 byte 的 SMEM 都不浪费。写入 sA 时也要用同一个 swizzle(row, col) 换算列号;XOR 只动列号的低 3 位,结果仍落在 0..15 之内,不会越界。

工业实现还有一条更彻底的路。CUTLASS 的 SIMT GEMM 把 A tile 转置着存:include/cutlass/gemm/threadblock/default_mma_core_simt.h:308 的 SmemLayoutA 是 ColumnMajor,相当于 sA[BK][BM],同一个 kk 的一整列在 SMEM 里连续;每线程用一条 LDS.128 取 4 个(同文件 :367,LaneM 取 128 位能装下的元素数与 thread tile 的较小者,FP32 即 4)。光转置还不够,线程映射也要跟着改:若每线程仍取相邻的 8 行,warp 内 8 个 float4 的起点相距 32 字节,仍有 2-way 冲突。CUTLASS 把 TM=8 拆成两段 4 个、两段之间隔开整个 warp 在 M 方向的跨度(include/cutlass/gemm/warp/mma_simt_tile_iterator.h:242 的 m * Policy::WarpShape::kRow),同一条指令里各 lane 的 16 字节块首尾相接,冲突才归零。第 12 章会看到 Tensor Core 路径上 CuTe 的 swizzle 做法。

消冲突的收益同样没有可引用的 H100 数字。值得一提的是 Boehm 的博文:他写过两版专门消 bank conflict 的 kernel,冲突确实消掉了,整体却比没消的版本慢,最终没有收进正文。布局改动会连带影响寄存器用量和访存指令的宽度,值不值要以 profile 为准。

11.6 优化 3:寄存器 Blocking 与读取顺序

最内层的 mul-add 循环:

for (int kk = 0; kk < BK; ++kk) {
    for (int i = 0; i < TM; ++i) a[i] = sA[...][kk];
    for (int j = 0; j < TN; ++j) b[j] = sB[kk][...];
    for (int i = 0; i < TM; ++i)
        for (int j = 0; j < TN; ++j)
            c[i][j] += a[i] * b[j];
}

这段代码每次 kk 都要重新 load a 和 b。如果 BK=16,每线程每个 K 步共 16 × (TM + TN) = 256 个标量的 SMEM 读取。b 的 8 个元素在 sB 的一行里连续,编译器已经把它们合成两条 LDS.128;a 的 8 个元素分在 sA 的 8 行,只能是 8 条标量 LDS——nvcc 13.4 为 sm_90a 编译 11.3 节的 kernel,每个 kk 是 8 条 LDS + 2 条 LDS.128 配 64 条 FFMA(不同 CUDA 版本可能不同)。

改进方式是把 kk 展开 4 步,先把 4 步的 a、b 一起读进寄存器,再集中做乘加:

for (int kk = 0; kk < BK; kk += 4) {
    float a[4][TM], b[4][TN];
    #pragma unroll
    for (int u = 0; u < 4; ++u) {
        for (int i = 0; i < TM; ++i) a[u][i] = sA[...][kk + u];
        for (int j = 0; j < TN; ++j) b[u][j] = sB[kk + u][...];
    }
    #pragma unroll
    for (int u = 0; u < 4; ++u) {
        #pragma unroll
        for (int i = 0; i < TM; ++i)
            #pragma unroll
            for (int j = 0; j < TN; ++j)
                c[i][j] += a[u][i] * b[u][j];
    }
}

读取的字节数一个没少,变的是指令形态:sA 按行存,同一行相邻 4 个 kk 在 SMEM 里连续,展开后编译器能把 a[0..3][i] 合成一条 LDS.128。nvcc 13.4 编译所得:每 4 个 kk 的 SMEM 读取从 40 条指令(32 条 LDS + 8 条 LDS.128)降到 16 条 LDS.128,代价是寄存器从 96 个涨到 128 个(每 SM 仍驻留 2 个 block;不同 CUDA 版本可能不同)。4 步的读取都排在乘加之前,在途的 SMEM 读取更多,延迟也更容易被掩盖。

c[i][j] += a[u][i] * b[u][j] 本身就是外积累加(outer product)的写法,CUTLASS 的 SIMT 路径在 thread 级也是这样组织乘加的(include/cutlass/gemm/thread/mma_sm50.h)。

11.7 性能演进表

把所有优化加上:

版本 优化 解决的问题 公开参照(Boehm,RTX A6000,占 cuBLAS 的比例)
v0 朴素 —— 8.5%
v1 Thread tile 4×4 访存指令减少四分之三 无对应版本
v2 + Block tile + SMEM(8×8 thread tile) 显式复用,block 级算术强度 32 68.7%(2D blocktiling)
v3 + Double buffer (cp.async) 加载与计算重叠 未实现
v4 + Bank conflict fix SMEM 读取不再串行化 消冲突版反而更慢,未收入正文
v5 + 寄存器 blocking + unroll SMEM 读取合并成 LDS.128 ——
参照 cuBLAS SGEMM —— 100%(23250 GFLOPs/s,约为 A6000 FP32 标称峰值的 60%)

关于这张表:本专栏没有条件在 H100 上逐版实测。右列是能找到的公开实测,GPU 是 A6000 而不是 H100,Boehm 的各版与本章各版也不是一一对应,只作量级参照。读者手上有卡的话,这正是最值得亲手复现的一张表。

Boehm 后面几版的数字还说明,手写 SIMT kernel 可以逼近 cuBLAS SGEMM:向量化访存 78.4%,参数自动调优 84.8%,warptiling 93.7%。要分清这里其实是两件不同的事:

  • 手写 kernel 与 cuBLAS SGEMM 的差距,与 Tensor Core 无关——cuBLAS SGEMM 走的也是 FP32 SIMT 的 FFMA 通路。差距来自更细的指令排布、更深的流水、按矩阵尺寸选 tile 的启发式,是"同一条赛道上的工艺差距"。
  • 真正的天花板是赛道本身:H100 SXM5 的 FP32 峰值只有 67 TFLOPs,而 FP16 Tensor Core 稠密峰值是 989 TFLOPs——差约 15 倍。就算把 SGEMM 打磨到 100%,也只是 67。

这就是为什么第 12 章必须引入 Tensor Core——不是 SIMT 的写法不够好,是这条赛道本身就短。

11.8 这一章的小结与下一章

Tiled GEMM 是 GEMM 优化的"地基":

  1. 三层 tile(block / warp / thread)让数据在 SMEM 和寄存器中层层复用。BM=BN=128、BK=16、FP32 时,每个 K 步搬 16 KB、算 524288 FLOPs,block 级算术强度 = 32 FLOPs/byte——比朴素 GEMM 按访存指令算的 0.25 高两个数量级,已经越过 H100 SXM5 FP32 通路的 Roofline 临界点(约 20)。
  2. Double buffer + cp.async 让 HBM 拷贝和计算重叠,隐藏 HBM 延迟。
  3. Bank conflict 处理(padding 或 swizzle)确保 SMEM 带宽。
  4. 寄存器 blocking 和 unroll 让编译器把 SMEM 读取合并成 LDS.128,减少 SMEM 访问指令。
  5. 手写 SIMT 可以逼近 cuBLAS SGEMM(Boehm 在 A6000 上做到了 93.7%),再往上是与 cuBLAS 比工艺;而整条 FP32 SIMT 赛道的峰值(H100 SXM5 67 TFLOPs)本身就只有 FP16 Tensor Core 稠密峰值的约 1/15。

第 12 章我们引入 Tensor Core——mma.sync 指令、ldmatrix 指令、layout swizzle。这是 GEMM 性能再上一个数量级的关键——换的是赛道,不是工艺。

本章动手练习:

  1. 把 v0..v5 都实现一遍,记录性能演进。
  2. 用 Nsight Compute 看 v2 vs v4 的 l1tex__data_bank_conflicts_pipe_lsu_mem_shared_op_ld.sum 指标(store 方向是 _op_st),验证 +1 padding 把冲突从 8-way 降到 2-way,swizzle 版则应降到 0。
  3. 思考:为什么 BM=BN=128 比 BM=BN=64 更优?(提示:按 11.8 节第 1 条的算法算两者的 block 级算术强度,与 FP32 临界点约 20 比较,再对照 SMEM 占用)