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

第 15 章 FA2 前向:Tiling 与 Online Softmax

作者 杨艺韬 · 5,893 字 · 发布于 · 更新于

FA2 相对 FA1 的改动不在数学,在并行度的重新分配——外层循环从 K/V 换成 Q。 这一章把这个改动落到一份 fragment 级的 forward kernel 骨架上。

15.1 FA2 vs FA1:parallelism 的重新分配

FlashAttention v2 论文(Dao, 2023)相比 v1 的核心改动不是新算法——两版的输出在数学上等价——而是四处工程改写。

改动 1:循环顺序对调。 FA1 论文 Algorithm 1 的外层循环是 K/V 块、内层是 Q 块;FA2 论文第 3.2 节写的是 swapping the order of the loop(outer loop over row blocks and inner loop over column blocks),改成外层 Q、内层 K/V,并注明这个想法和下面的按序列长度并行最早由 Phil Tillet 在 Triton 实现里提出。论文给的动机是并行度:外层的 Q 块之间互不依赖,可以分给不同的 thread block(见改动 2)。顺带的收益是 O 和 m,ℓm,\ell 全程留在片上,只在外层循环结束时写一次 HBM,而不是每处理一个 K/V 块就把它们读回来再写出去。第 14 章 14.7 节列过这张对照表。

改动 2:多一条并行轴。 FA1 的 grid 只在 (batch × head) 上铺开——一个 thread block 负责一个 head 的全部 attention。batch × head 小于 SM 数时(长上下文推理、小 batch 训练都会遇到),大量 SM 直接闲置。循环顺序换成外层 Q 之后,每个 Q 块天然是一个独立的 thread block,于是 grid 变成 (Q 块数 × head 数 × batch)——序列长度本身成了一条并行轴,长序列反而更容易填满 GPU。

改动 3:warp 间的工作划分从 split-K 换成 split-Q。 FA1 把 K/V 块按列切给 block 内的各个 warp("split-K"),每个 warp 算出 S 的一部分列,因此每一轮都要把各 warp 的部分结果写进 SMEM 再规约。FA2 反过来,把 Q 块按行切给各 warp("split-Q"),每个 warp 独占若干行、独立完成整行的 online softmax 与 O 累加——warp 之间在主循环里不再需要通信。这是本章 15.3 节骨架里 Br_warp = Br / N_WARPS 那一行的由来。

改动 4:削掉非矩阵乘的 op。 Tensor Core 算矩阵乘极快,但 exp、除法、rescale 走的是普通 SIMT 单元(exp 还要过 SFU)。FA2 论文第 3.1 节给的 A100 口径是 FP16/BF16 矩阵乘 312 TFLOPs/s、非矩阵乘 FP32 只有 19.5 TFLOPs/s,每个非矩阵乘 FLOP 贵 16 倍。FA2 把每一轮的 O /= l 推迟到最后统一做一次,中间只留一次 rescale;给反向存的统计量也从 mm、ℓ\ell 两个缩成一个 logsumexp L=m+log⁡ℓL = m + \log \ell。

FA2 论文摘要给出的整体效果是相对 FA1 约 2×,在 A100 上达到理论峰值的 50–73%。

15.2 Tile 大小的选择

FA2 forward 的 tile 大小有几个关键参数:

Br  : Q block 的行数 (典型 64 或 128)
Bc  : K/V block 的列数 (典型 64)
d   : head dim (典型 64, 128, 256)

约束:

  1. SMEM 容量:要放下 Q_block (Br × d) + K_block (Bc × d) + V_block (Bc × d)。FP16 时 = 2 × (Br + 2 × Bc) × d 字节。
  2. 寄存器:一个 warp 的 O 累加器是 (Br/N_warps) × d 个 fp32,摊到 32 个 lane 上,每线程 (Br/N_warps) × d / 32 个。Br=128、N_warps=4、d=64 时是 64 个 fp32——已经占掉每线程 255 个寄存器上限的四分之一。
  3. 算术强度:tile 越大,每次 K_tile 的 GEMM 越大,Tensor Core 利用率越好。

几档候选配置(FP16, d=64)的 SMEM 账:

Br Bc SMEM 占用 = 2·(Br+2·Bc)·d Block 内 warp 数
64 64 24 KB 4 (128 thread)
128 64 32 KB 4
128 128 48 KB 4

FA2 论文第 3.3 节说块大小一般在 {64, 128} × {64, 128} 里按 head dim 和 SMEM 容量手工挑。pin 版本的官方实现里,d=64 前向(无 dropout)选的是 Br=128、Bc=128、4 个 warp,旁边的注释记着 8 个 warp 在 seqlen=2k 时慢 18%(flash-attn/csrc/flash_attn/src/flash_fwd_launch_template.h:207、:210);d=128 在 A100/H100 上选 Br=128、Bc=64、4 个 warp(同文件 :262)。

不同 head dim 也要调整。d=128 时所有 SMEM 翻倍,需要降低 Bc 维持 SMEM 预算。

15.3 完整的 FA2 Forward Kernel 骨架

下面是 FA2 forward kernel 的骨架。cp_async_*、ldmatrix_*、mma_m16n8k16 这几个辅助函数没有展开(各是几行 inline PTX,写法见第 12 章 12.2、12.3 节;cp.async 见第 11 章),另外简化了 swizzle 细节和一部分边界处理,但结构反映工业级实现:

template <int Br, int Bc, int d, int N_WARPS = 4>
__global__ void flash_attn_fwd(
    const half* __restrict__ Q,    // [B, H, N, d]
    const half* __restrict__ K,    // [B, H, N, d]
    const half* __restrict__ V,    // [B, H, N, d]
          half* __restrict__ O,    // [B, H, N, d]
          float* __restrict__ LSE, // [B, H, N], for backward
    int N,
    int H,                         // head 数,下面算 q_offset 要用(少了它编不过)
    float softmax_scale,           // 1 / sqrt(d)
    bool is_causal
) {
    // Grid: (ceil(N / Br), H, B)
    // Block: 128 threads (4 warps)
    int q_tile_idx = blockIdx.x;
    int head_idx = blockIdx.y;
    int batch_idx = blockIdx.z;

    int tid = threadIdx.x;
    int warp_id = tid / 32;
    int lane_id = tid % 32;

    // ============ SMEM allocation ============
    extern __shared__ half smem[];
    half* sQ = smem;                              // [Br, d]
    half* sK = sQ + Br * d;                       // [Bc, d]
    half* sV = sK + Bc * d;                       // [Bc, d]

    // ============ Pointer offset for this (B, H) ============
    int q_offset = ((batch_idx * H + head_idx) * N + q_tile_idx * Br) * d;
    int kv_offset = (batch_idx * H + head_idx) * N * d;
    int o_offset = q_offset;
    int lse_offset = (batch_idx * H + head_idx) * N + q_tile_idx * Br;

    // ============ Load Q tile to SMEM (one-time) ============
    // 约定: 两个 load 辅助函数把第 N 行之后的越界行填 0 (cp.async 的 src-size 补零)
    cp_async_load_q_tile(sQ, Q + q_offset);
    cp_async_commit_and_wait();
    __syncthreads();

    // ============ Output accumulator (in registers) ============
    // Each warp handles Br/N_WARPS rows of Q.
    // Each warp keeps Br_warp × d output accumulator in fp32 fragments.
    constexpr int Br_warp = Br / N_WARPS;
    constexpr int MMAS_M = Br_warp / 16;     // # of mma.m16n8k16 in M
    constexpr int MMAS_D = d / 8;             // # of mma in N (output dim)
    float O_acc[MMAS_M][MMAS_D][4] = {0};     // 累加 fragment

    // 注意: 千万别写成 float row_max[MMAS_M][2] = {-INFINITY};
    // C++ 聚合初始化只会把首元素设成 -INFINITY, 其余元素被值初始化为 0 ——
    // 那些行的 running max 一开局就是 0, 不再是真正的行最大值; 某行分数整体
    // 远低于 0 时 expf(s - 0) 在 fp32 下下溢成 0, 结果悄悄算错。
    float row_max[MMAS_M][2];                 // 每个 m16 块里, lane 持 2 行 (l/4 和 l/4+8)
    #pragma unroll
    for (int i = 0; i < MMAS_M; ++i)
        row_max[i][0] = row_max[i][1] = -INFINITY;

    float row_sum[MMAS_M][2] = {0};           // 这个 {0} 是对的: 值初始化就是全 0

    // ============ Loop over K tiles ============
    int k_tile_end = is_causal
        ? min(N, (q_tile_idx + 1) * Br)
        : N;
    for (int k_tile = 0; k_tile < k_tile_end; k_tile += Bc) {
        // ---------- Load K, V tile to SMEM ----------
        cp_async_load_kv_tile(sK, K + kv_offset + k_tile * d);
        cp_async_load_kv_tile(sV, V + kv_offset + k_tile * d);
        cp_async_commit_and_wait();
        __syncthreads();

        // ---------- 1) S_block = Q @ K^T (using mma.sync) ----------
        // S_block is [Br_warp, Bc] in fragments (per warp).
        constexpr int MMAS_N = Bc / 8;
        float S_frag[MMAS_M][MMAS_N][4] = {0};

        for (int kk = 0; kk < d; kk += 16) {
            // ldmatrix Q fragments
            unsigned q_frag[MMAS_M][4];
            for (int i = 0; i < MMAS_M; ++i) {
                int row = warp_id * Br_warp + i * 16;
                ldmatrix_x4(sQ, row, kk, q_frag[i]);
            }
            // ldmatrix K fragments: 不要 .trans —— sK 按 [Bc, d] 存, d 维连续,
            // 对 B = K^T 来说正是 mma .col 要的"沿 K 维相邻" (第 12 章 12.3 节)
            unsigned k_frag[MMAS_N][2];
            for (int j = 0; j < MMAS_N; ++j) {
                int col = j * 8;
                ldmatrix_x2(sK, col, kk, k_frag[j]);
            }
            // mma accumulate into S_frag
            for (int i = 0; i < MMAS_M; ++i)
                for (int j = 0; j < MMAS_N; ++j)
                    mma_m16n8k16(S_frag[i][j], q_frag[i], k_frag[j]);
        }

        // ---------- 2) Apply softmax_scale and causal mask ----------
        for (int i = 0; i < MMAS_M; ++i) {
            for (int j = 0; j < MMAS_N; ++j) {
                #pragma unroll
                for (int e = 0; e < 4; ++e) {
                    S_frag[i][j][e] *= softmax_scale;
                    // Causal mask: lane_id 0/4/8/12... 持有不同列, 需精确计算
                    int my_row = q_tile_idx * Br + warp_id * Br_warp
                                 + i * 16 + (lane_id / 4) + (e / 2) * 8;
                    int my_col = k_tile + j * 8 + (lane_id % 4) * 2 + (e % 2);
                    // my_col >= N: N 不是 Bc 整数倍时, 最后一个 tile 的补零列也要 mask
                    if ((is_causal && my_col > my_row) || my_col >= N) {
                        S_frag[i][j][e] = -INFINITY;
                    }
                }
            }
        }

        // ---------- 3) Online softmax: update row_max, row_sum ----------
        // Each warp computes max/sum across columns within the tile.
        // S_frag layout: lane (l/4, l%4) holds (rows [l/4, l/4+8], cols [2*(l%4), 2*(l%4)+1]).
        // We need: for each row, find max and sum across all Bc columns.
        for (int i = 0; i < MMAS_M; ++i) {
            // 每 mma 块持有 16 行, 每 lane 拥有 2 行 (l/4 and l/4+8)
            // 拿到 row 内所有 col 的 max
            for (int row_local = 0; row_local < 2; ++row_local) {
                float m_block = -INFINITY;
                for (int j = 0; j < MMAS_N; ++j) {
                    for (int col_local = 0; col_local < 2; ++col_local) {
                        m_block = fmaxf(m_block,
                                  S_frag[i][j][row_local * 2 + col_local]);
                    }
                }
                // Warp 内 reduce: 每行的 max 分布在 4 个 lane 上 (相同 lane_id/4)
                // 用 shfl_xor 在 4 个 lane 之间归约
                m_block = fmaxf(m_block, __shfl_xor_sync(0xFFFFFFFF, m_block, 1));
                m_block = fmaxf(m_block, __shfl_xor_sync(0xFFFFFFFF, m_block, 2));

                float m_old = row_max[i][row_local];
                float m_new = fmaxf(m_old, m_block);
                // 到目前为止整行都被 mask 时 m_new 仍是 -INF, 直接减会得到
                // expf(-INF - (-INF)) = NaN; 按官方 softmax.h 的 Check_inf 改用 0 作基准
                float m_use = (m_new == -INFINITY) ? 0.0f : m_new;
                float alpha = expf(m_old - m_use);

                // 更新 P_frag (in-place 改写 S_frag, 同时算 row sum 增量)
                float l_inc = 0.0f;
                for (int j = 0; j < MMAS_N; ++j) {
                    for (int col_local = 0; col_local < 2; ++col_local) {
                        float p = expf(
                            S_frag[i][j][row_local * 2 + col_local] - m_use);
                        S_frag[i][j][row_local * 2 + col_local] = p;
                        l_inc += p;
                    }
                }
                // Warp 归约 l_inc
                l_inc += __shfl_xor_sync(0xFFFFFFFF, l_inc, 1);
                l_inc += __shfl_xor_sync(0xFFFFFFFF, l_inc, 2);

                row_sum[i][row_local] = row_sum[i][row_local] * alpha + l_inc;
                row_max[i][row_local] = m_new;

                // ---------- 4) 缩放之前累积的 O_acc ----------
                for (int j_d = 0; j_d < MMAS_D; ++j_d) {
                    for (int e = 0; e < 4; ++e) {
                        // 只有 row_local 对应的元素需要 alpha 缩放
                        if ((e / 2) == row_local) {
                            O_acc[i][j_d][e] *= alpha;
                        }
                    }
                }
            }
        }

        // ---------- 5) O_acc += P @ V ----------
        // P 是 fp32 fragment, 要先 cast 回 fp16 给 mma 用
        // 实际上 FA2 选择 fp16 P × fp16 V, fp32 累加
        // P_frag (fp16) layout: [Br_warp, Bc] = MMAS_M × MMAS_N × 4
        // 每个 (i, j) 的 4 个 fp32 打包成 **2 个** unsigned(每个装 2 个 half),
        // 所以第二维是 MMAS_N * 2 而不是 MMAS_N。
        unsigned P_frag_fp16[MMAS_M][MMAS_N * 2];
        for (int i = 0; i < MMAS_M; ++i) {
            for (int j = 0; j < MMAS_N; ++j) {
                // 把 4 个 fp32 转成 4 个 fp16, 按行两两打包: (c0,c1) 行 l/4, (c2,c3) 行 l/4+8
                __half2 p01, p23;
                p01.x = __float2half(S_frag[i][j][0]);
                p01.y = __float2half(S_frag[i][j][1]);
                p23.x = __float2half(S_frag[i][j][2]);
                p23.y = __float2half(S_frag[i][j][3]);
                P_frag_fp16[i][j * 2 + 0] = *reinterpret_cast<unsigned*>(&p01);
                P_frag_fp16[i][j * 2 + 1] = *reinterpret_cast<unsigned*>(&p23);
            }
        }

        // mma: O_acc[i][d_j] += P_frag[i][j] @ V_frag[j][d_j]
        for (int j = 0; j < MMAS_N; j += 2) {  // P fragment 一组 16 列
            // V 要 .trans: sV 按 [Bc, d] 存, d 维 (即 mma 的 N 维) 连续,
            // 而 .col 要每个 lane 拿到沿 Bc 维 (mma 的 K 维) 相邻的两个元素
            unsigned v_frag[MMAS_D][2];
            for (int d_j = 0; d_j < MMAS_D; ++d_j) {
                int v_row = j * 8;
                int v_col = d_j * 8;
                ldmatrix_x2_trans(sV, v_row, v_col, v_frag[d_j]);
            }
            for (int i = 0; i < MMAS_M; ++i) {
                for (int d_j = 0; d_j < MMAS_D; ++d_j) {
                    // mma 的 A 要 4 个寄存器: 8 列一组的 j 和 j+1 各贡献 2 个
                    // (每组里前一个是行 l/4, 后一个是行 l/4+8), 合起来正好 16 列
                    unsigned p_input[4] = {
                        P_frag_fp16[i][j * 2 + 0], P_frag_fp16[i][j * 2 + 1],
                        P_frag_fp16[i][j * 2 + 2], P_frag_fp16[i][j * 2 + 3]
                    };
                    mma_m16n8k16(O_acc[i][d_j], p_input, v_frag[d_j]);
                }
            }
        }
        __syncthreads();
    }

    // ============ Final normalization: O = O_acc / row_sum ============
    for (int i = 0; i < MMAS_M; ++i) {
        for (int row_local = 0; row_local < 2; ++row_local) {
            float l = row_sum[i][row_local];
            float scale = (l == 0.0f) ? 1.0f : 1.0f / l;  // 整行被 mask: O 保持 0
            for (int d_j = 0; d_j < MMAS_D; ++d_j) {
                for (int e = 0; e < 4; ++e) {
                    if ((e / 2) == row_local) {
                        O_acc[i][d_j][e] *= scale;
                    }
                }
            }
        }
    }

    // ============ Write O to HBM ============
    // o_offset 已含 q_tile_idx * Br 行, 这里只用 tile 内行号 r
    // (直接从寄存器写, 未经 SMEM 合并; 官方实现先写 SMEM 再合并写回)
    for (int i = 0; i < MMAS_M; ++i) {
        for (int row_local = 0; row_local < 2; ++row_local) {
            int r = warp_id * Br_warp + i * 16 + (lane_id / 4) + row_local * 8;
            if (q_tile_idx * Br + r >= N) continue;   // 最后一个 Q tile 的越界行不写
            for (int d_j = 0; d_j < MMAS_D; ++d_j) {
                int col = d_j * 8 + (lane_id % 4) * 2;
                half2 v;
                v.x = __float2half(O_acc[i][d_j][row_local * 2 + 0]);
                v.y = __float2half(O_acc[i][d_j][row_local * 2 + 1]);
                *reinterpret_cast<half2*>(&O[o_offset + r * d + col]) = v;
            }
        }
    }

    // ============ Write LSE for backward ============
    // 每个 warp 写自己那 Br_warp 行; 同一行的 4 个 lane 值相同, 由 lane%4==0 写
    for (int i = 0; i < MMAS_M; ++i) {
        for (int row_local = 0; row_local < 2; ++row_local) {
            int r = warp_id * Br_warp + i * 16 + (lane_id / 4) + row_local * 8;
            if (lane_id % 4 == 0 && q_tile_idx * Br + r < N) {
                float l = row_sum[i][row_local];
                LSE[lse_offset + r] = (l == 0.0f)
                    ? INFINITY                          // 同官方 softmax.h: 整行 mask
                    : row_max[i][row_local] + logf(l);
            }
        }
    }
}

注意:上面是简化骨架,省略了 swizzle layout、cp.async pipeline 和访存合并;边界只处理到「越界行由 load 补零、越界列 mask 成 -INF、越界行不写回」。三处 -INF 防护(m_use、l == 0 两处)照官方 flash-attn/csrc/flash_attn/src/softmax.h:154-156(缩放因子)、:76(求 exp 时的行最大值)、:179-180(归一化与 LSE)的写法:本骨架自带的 causal mask 下每行第一个 K tile 至少有第 0 列可见,不会出现 m_new == -INF;但读者一旦加上滑动窗口、padding 之类的 mask,某行在前几个 tile 里全被 mask 就会得到 expf(-INF - (-INF)) = NaN,并一路传进 O。骨架里 row_max 那几行也是照抄不会错的写法——它避开了 C++ 聚合初始化的一个经典陷阱:float a[N][2] = {-INFINITY}; 只把 a[0][0] 设成 -INFINITY,其余元素一律值初始化为 0。running max 从 0 起步时,softmax 的平移不变性让普通分数下结果仍然对,所以这个 bug 很难被小测试抓到;可一旦某行分数整体远低于 0,expf(s - 0) 在 fp32 下下溢成 0,结果不会崩、只会悄悄算错。完整实现请参考 flash-attn/csrc/flash_attn/src/flash_fwd_kernel.h(GitHub 上的同一份)。

15.4 关键实现要点

15.4.1 Online Softmax 的 fragment 级实现

第 6 章的 online softmax 是每个线程顺序扫一段元素、再做 warp/block 归约的写法。落到 fragment 级实现,关键挑战是 mma 的 fragment layout 让"per-row"操作变复杂——每个 lane 持有的不是连续的"一行",而是分散的"两行的几列"。

具体做法:

  1. 每行 max:每 lane 在自己持有的列上算 local max,然后用 shfl_xor 在持有同一行的 4 个 lane 之间归约。
  2. 每行 sum:同上,归约 sum。
  3. 同步更新 m, l, O:所有更新都基于 fragment 内的本地数据。行最大值仍是 -INF(到目前为止整行被 mask)时要换 0 作基准,否则得到 NaN(第 6 章 6.4 节同一个问题)。

15.4.2 P 矩阵的 fp16 / fp32 切换

S_frag 是 fp32 累加(mma 输出),但 P @ V 的 mma 输入要求 fp16。所以中间需要把 fp32 P_frag 转回 fp16。

好在这一步不用过 SMEM:m16n8k16 的 C fragment 里,相邻两个 n8 块(16 列)在每个 lane 上的分布,正好就是下一次 mma 的 A fragment 布局,所以转成 fp16 后在寄存器里两两打包就能直接喂给 P @ V(即 15.3 节骨架里名为 p_input 的那四个寄存器)。官方实现也是这么做的:flash-attn/csrc/flash_attn/src/flash_fwd_kernel.h:416 用 convert_type 转精度,:434 用 convert_layout_acc_Aregs 把同一块寄存器重新解释成 A 操作数,再交给 gemm_rs(A 在寄存器、B 从 SMEM 读)。精度方面,P 的元素都在 [0, 1] 里,转 fp16 只是一次舍入,O 的累加仍是 fp32。

15.4.3 Causal Mask 的位置

Causal mask 必须在 softmax 之前 apply(对 -inf 取 exp 是 0,不影响 sum)。在 fragment 级,每个 lane 知道自己持有哪些 (row, col),直接对越界元素写 -inf。N 不是 Bc 整数倍时,最后一个 K tile 里补零的列也要这样 mask,否则 0 分数会以 e0−me^{0 - m} 的权重混进 sum。

15.4.4 Swizzle Layout

Q/K/V 在 SMEM 里的布局要做 swizzle,否则 ldmatrix 触发 bank conflict(第 12 章 12.4 节)。官方 FA2 用的是 CuTe 的 Swizzle,参数由 head dim 决定(flash-attn/csrc/flash_attn/src/kernel_traits.h:70、:72、:80):

static constexpr int kBlockKSmem = kHeadDim % 64 == 0 ? 64 : 32;
static constexpr int kSwizzle    = kBlockKSmem == 32 ? 2 : 3;
// ...
composition(Swizzle<kSwizzle, 3, 3>{},
            Layout<Shape<_8, Int<kBlockKSmem>>,
                   Stride<Int<kBlockKSmem>, _1>>{})

也就是 head dim 是 64 的整数倍时用 Swizzle<3,3,3>(SMEM 里 K 方向一次铺 64 个元素),否则退到 Swizzle<2,3,3>(铺 32 个)。源码注释还特意说明了这里必须用 kBlockKSmem 而不是 kHeadDim,否则 d=128 会算出错误结果(flash-attn/csrc/flash_attn/src/kernel_traits.h:81 的原文就是 "This has to be kBlockKSmem, using kHeadDim gives wrong results for d=128");同文件 :113 还记了另一笔账——d=128 时用 kBlockKSmem 比 kBlockKGmem 快 6~10%,同样是 bank conflict 的缘故。注意这里的 Swizzle<3,3,3> 第二个参数是 3,因为它作用在元素偏移上(23=82^3 = 8 个 half = 16 字节);第 13 章引的 Layout_*_SW128_Atom_Bits 写的是 Swizzle<3,4,3>,因为那一族的 swizzle 挂在 smem_ptr_flag 上,作用在 SMEM 的字节地址上(242^4 字节 = 16 字节;名字里的 Bits 指的是内层 layout 按比特计,见 cutlass-4.7.0/include/cute/pointer_swizzle.hpp:87 与 cutlass-4.7.0/include/cute/pointer_flagged.hpp:113)。两者说的是同一档 128B swizzle。第 13 章讲过的概念在这里直接落地。

15.5 性能调优要点

15.5.1 K 维度的 cp.async pipeline

外层 K 循环里,加载下一个 K_block + V_block 应该和当前 block 的计算重叠:

// 启动 K_tile 0 的加载
cp_async_load(sK[0], sV[0], 0);
cp_async_commit();

for (int k_tile = 0; k_tile < N; k_tile += Bc) {
    int next = (k_tile / Bc + 1) % 2;
    int cur = (k_tile / Bc) % 2;

    // 启动下一个 K_tile 加载
    // (写的是 next 缓冲区, 它上一轮刚被算完, 靠循环末尾的 __syncthreads 保证)
    if (k_tile + Bc < N) {
        cp_async_load(sK[next], sV[next], k_tile + Bc);
        cp_async_commit();
        cp_async_wait_group<1>();   // 只等当前这组, 刚发出的下一组留在路上
    } else {
        cp_async_wait_group<0>();   // 最后一轮没有新组, 全部等完
    }
    __syncthreads();

    // 算当前 K_tile
    compute(sK[cur], sV[cur], &O_acc, ...);
    __syncthreads();                // 所有 warp 用完 cur, 下一轮才能往里写
}

这样 SM 的 Tensor Core 在算当前 tile 时,下一个 tile 的 HBM 加载在后台进行。两处容易写错:等待必须是 wait_group 1(wait_group 0 会把刚发出的预取也一起等掉,重叠就没了);循环末尾的 __syncthreads() 不能省,否则快的 warp 会在慢的 warp 还在读 cur 时往同一块缓冲区预取。官方 FA2 没有用 K/V 各两份缓冲,而是在一份 K、一份 V 之间交错:算 QK⊤QK^\top 时 V 在加载,算 softmax 和 PVPV 时下一块 K 在加载(flash-attn/csrc/flash_attn/src/flash_fwd_kernel.h:385-435),K/V 部分的 SMEM 省一半。

15.5.2 寄存器压力管理

FA2 forward 的寄存器需求很大。按 15.3 节骨架、Br=128、Bc=64、d=64、4 个 warp 数一下(MMAS_M=2、MMAS_N=8、MMAS_D=8):

  • 输出累加器 O_acc:MMAS_M × MMAS_D × 4 = 64 个 fp32
  • S/P fragment:S_frag 是 MMAS_M × MMAS_N × 4 = 64 个 fp32,转成 fp16 的 P_frag 再占 32 个
  • Q/K/V fragment:q_frag 8 个、k_frag 16 个、v_frag 16 个
  • m、l 各 MMAS_M × 2 = 4 个,外加 alpha、地址、循环变量等临时量

三项粗加就接近 200 个。把骨架配上最简单的辅助函数实现(行优先 SMEM、无 swizzle)用 nvcc 13.4 编译,-Xptxas -v 报这一档每线程 233 个寄存器(sm_80)/ 254 个(sm_90a),0 spill;d=128(Br=128、Bc=64)已顶到 255 个并开始 spill;Br=Bc=128 时则出现 800 字节的 stack frame——骨架里几处循环没有展开,动态下标的 fragment 数组被放进了 local memory;给所有内层循环补上 #pragma unroll 后 stack frame 降到 150 多字节,剩下的是寄存器确实不够造成的 spill(不同 CUDA 版本结果可能不同)。官方 FA2 在 d=64 用的正是 128×128(见 15.2 节),它的 fragment 调度比这份骨架省寄存器得多。硬件的天花板是每线程 255 个(第 4 章 §4.2.3);为了 occupancy 主动压配额也会把你推下悬崖——写 __launch_bounds__(128, 4) 要求每 SM 驻 4 个 block,编译器就得把每线程压到 128 个以内,上面这份骨架必然大量 spill。

减少压力的技巧:

  • fragment 数组的下标必须是编译期常量:对 S_frag、O_acc 这类数组的循环都要 #pragma unroll(循环边界本身是模板常量),否则数组会被放进 local memory。
  • 别把 m、l 降精度——它们是 online softmax 的 running max / running sum,l 最大可到序列长度,超过 65504 就溢出 fp16,而且 fp16 只有 11 位有效精度,累加成千上万个小数误差很大,降精度是拿正确性换寄存器,不划算。该省的是量大的那部分:S_frag 算完立刻转成 fp16 的 P_frag,让 fp32 的 S_frag 尽早死掉。
  • 用 __launch_bounds__ 或 -maxrregcount 显式给编译器一个配额,让它自己权衡,而不是让它猜;配额定多少要配合 -Xptxas -v 看 spill 是否为 0(第 4 章 §4.2.3)。
  • 减少 unroll 程度(牺牲 ILP 换寄存器)。

15.5.3 Block Size 与 Occupancy

128 thread/block × 4 active block/SM = 512 active threads/SM。这只是 SM 上限 2048 的 1/4——occupancy 约 25%。按 15.5.2 节的 233 个寄存器/线程算还更低:分配时按 8 的倍数取整到 240,一个 block 要 128 × 240 = 30720 个寄存器,每 SM 的 65536 个只够驻 2 个 block,occupancy 12.5%。

但 FA2 这类 kernel 并不靠堆驻留 warp 来隐藏延迟,而是靠每个 warp 内部的 ILP(一次发出多条相互独立的 mma、ldmatrix)和 cp.async 预取。低 occupancy 的高 ILP kernel 可以比高 occupancy 的低 ILP kernel 快——这在 Tensor Core GEMM 里是常态,不是 FA 独有。注意这里说的是每 SM 的 occupancy,和 15.1 节「SM 闲置」说的 grid 级并行不是一回事:block 总数少于 SM 数时,ILP 再高也救不了空着的 SM。

15.6 性能:能期待到什么位置

本专栏没有 GPU 可跑 benchmark,这里只给有出处的锚点和随序列长度变化的趋势,不编造具体数字。

锚点来自 FA3 论文(Shah et al., 2024)在 H100 上对 FA2 的测量:FP16 下约 35% 的 Tensor Core 利用率(H100 SXM5 FP16/BF16 稠密峰值 989 TFLOPS,折算约 350 TFLOPs)。FA2 论文自己的口径与之一致:同一份实现不用 TMA、第四代 Tensor Core 等新特性直接跑在 H100 上,最高 335 TFLOPs/s(第 4.1 节)。作为对照,A100 上它能到理论峰值的 73%,而 FA3 在 H100 上 FP16/BF16 约 740 TFLOPs(75%,第 17 章 17.8 节)。这是官方 FA2 实现的水位——本节这份手写骨架省掉了 swizzle 细节、pipeline 深度调优和边界处理,只会更低。

趋势可以对照 FA2 论文 Fig. 5(A100 前向,seqlen 512 到 16k):

  • 序列越长,利用率越高。N 小时(短 prompt)主循环轮数少,Q tile 的加载、epilogue 的归一化和写回摊不开,非矩阵乘部分占比高;N 增大后主循环里的两个 GEMM 占绝对主导。Fig. 5 里这一趋势在 causal 情形最明显,非 causal 情形 1k 以后基本持平。
  • causal mask 会砍掉大约一半的有效工作,但被 mask 掉的 tile 如果不跳过就是纯浪费——所以 k_tile_end 那一行的裁剪在长序列上收益很大。
  • head dim 变大反而更容易跑高:d 越大,每个 tile 的 GEMM 越"方",Tensor Core 的 K 维流水越满。Fig. 5 里 d=128 的前向吞吐整体高于 d=64。

从手写骨架到官方实现,这个 gap 主要差在:

  • CUTLASS 级别的细致 fragment 调度
  • Hopper 的 TMA + WGMMA(FA3 的内容,第 17 章)
  • PTX 微优化(指令排布、依赖距离)

15.7 这一章的小结与下一章

第 15 章我们写出了一个结构完整的 FA2 forward kernel 骨架:

  1. FA2 vs FA1 的关键差别是循环顺序与并行划分:外层 Q 内层 K/V(FA1 相反)、序列长度成为并行轴、warp 间从 split-K 改成 split-Q、非矩阵乘 op 削到最少。
  2. 完整的 FA2 forward kernel 由 5 个阶段组成:load Q tile → loop K tile (load → S=QK^T → online softmax → O+=PV) → final normalize → write O/LSE。
  3. online softmax 在 fragment 级实现需要小心处理 lane 持有的不连续 row/col 布局,并对整行被 mask 的 -INF 做防护;K 用 ldmatrix、V 用 ldmatrix.trans,P 在寄存器里直接转成下一次 mma 的 A 操作数。
  4. cp.async pipeline 和寄存器压力管理是性能调优的关键。
  5. 官方 FA2 在 H100 上的水位是 ~35% Tensor Core 利用率(FA3 论文口径)——手写骨架能靠近它就已经是不小的成就,要再往上走必须换 Hopper 原生工具。

第 16 章我们处理 FA2 的反向。反向比前向复杂得多——需要重计算 S 和 P,需要原子写 dQ,且循环方向变成"外层 K,内层 Q"。读完第 16 章读者会拥有完整的 FA2 训练能力。

本章动手练习:

  1. 把上面的 FA2 forward 骨架完整实现(包括 swizzle layout 和正确的 lane→location 映射),扫 N=512/2048/4096/8192 测性能,验证 15.6 节说的"序列越长利用率越高"。
  2. 阅读 flash-attn/csrc/flash_attn/src/flash_fwd_kernel.h,对照 15.3 节的骨架找差异点(你会发现 swizzle、async、unroll 调度都更精细)。
  3. 思考:如果 head_dim=128,tile 大小该怎么调?SMEM 还放得下吗?