CUDA 算子工程:手写 FlashAttention v2 之路
第 15 章 FA2 前向:Tiling 与 Online Softmax
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 和 全程留在片上,只在外层循环结束时写一次 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;给反向存的统计量也从 、 两个缩成一个 logsumexp 。
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)
约束:
- SMEM 容量:要放下 Q_block (Br × d) + K_block (Bc × d) + V_block (Bc × d)。FP16 时 = 2 × (Br + 2 × Bc) × d 字节。
- 寄存器:一个 warp 的 O 累加器是 (Br/N_warps) × d 个 fp32,摊到 32 个 lane 上,每线程 (Br/N_warps) × d / 32 个。Br=128、N_warps=4、d=64 时是 64 个 fp32——已经占掉每线程 255 个寄存器上限的四分之一。
- 算术强度: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 持有的不是连续的"一行",而是分散的"两行的几列"。
具体做法:
- 每行 max:每 lane 在自己持有的列上算 local max,然后用
shfl_xor在持有同一行的 4 个 lane 之间归约。 - 每行 sum:同上,归约 sum。
- 同步更新 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 分数会以 的权重混进 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,因为它作用在元素偏移上( 个 half = 16 字节);第 13 章引的 Layout_*_SW128_Atom_Bits 写的是 Swizzle<3,4,3>,因为那一族的 swizzle 挂在 smem_ptr_flag 上,作用在 SMEM 的字节地址上( 字节 = 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 之间交错:算 时 V 在加载,算 softmax 和 时下一块 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_frag8 个、k_frag16 个、v_frag16 个 - 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 骨架:
- FA2 vs FA1 的关键差别是循环顺序与并行划分:外层 Q 内层 K/V(FA1 相反)、序列长度成为并行轴、warp 间从 split-K 改成 split-Q、非矩阵乘 op 削到最少。
- 完整的 FA2 forward kernel 由 5 个阶段组成:load Q tile → loop K tile (load → S=QK^T → online softmax → O+=PV) → final normalize → write O/LSE。
- online softmax 在 fragment 级实现需要小心处理 lane 持有的不连续 row/col 布局,并对整行被 mask 的
-INF做防护;K 用ldmatrix、V 用ldmatrix.trans,P 在寄存器里直接转成下一次 mma 的 A 操作数。 - cp.async pipeline 和寄存器压力管理是性能调优的关键。
- 官方 FA2 在 H100 上的水位是 ~35% Tensor Core 利用率(FA3 论文口径)——手写骨架能靠近它就已经是不小的成就,要再往上走必须换 Hopper 原生工具。
第 16 章我们处理 FA2 的反向。反向比前向复杂得多——需要重计算 S 和 P,需要原子写 dQ,且循环方向变成"外层 K,内层 Q"。读完第 16 章读者会拥有完整的 FA2 训练能力。
本章动手练习:
- 把上面的 FA2 forward 骨架完整实现(包括 swizzle layout 和正确的 lane→location 映射),扫 N=512/2048/4096/8192 测性能,验证 15.6 节说的"序列越长利用率越高"。
- 阅读
flash-attn/csrc/flash_attn/src/flash_fwd_kernel.h,对照 15.3 节的骨架找差异点(你会发现 swizzle、async、unroll 调度都更精细)。- 思考:如果 head_dim=128,tile 大小该怎么调?SMEM 还放得下吗?