CUDA 算子工程:手写 FlashAttention v2 之路
第 16 章 FA2 反向:dQ/dK/dV 的重计算
反向不是把前向倒着写一遍:前向用着顺手的数据布局,到反向几乎每一处都要重新安排。 这一章讲 dQ/dK/dV 三条梯度各自的并行策略,以及为什么 S/P 必须重算而不是存下来。
16.1 反向的数学:从 dO 到 dQ/dK/dV
设 attention 前向:
给定 ,我们要算:
链式法则一步步:
dV 最直接:
dP:
dS(softmax 的反向): 设 ,则 。所以:
定义 ,则:
dQ, dK:
总结,反向需要算的量:
1. D = rowsum(P ⊙ dP) = rowsum(P ⊙ (dO @ V^T))
2. dV = P^T @ dO
3. dS = P ⊙ (dP - D)
4. dK = dS^T @ Q
5. dQ = dS @ K
注意 这一项——它是 按行求和。论文里有个简化技巧:
注意这里的 是已归一化的 softmax 输出, 正是前向的输出行——所以这是恒等式,不带任何额外系数。
更简洁:。这避免了显式算 P,只需要前向已经写出去的 O 和上游传下来的 dO。
16.2 反向的循环顺序:外层 K,内层 Q
FA2 前向用的是"外层 Q,内层 K/V"(第 14 章 14.5 节;FA1 的前向恰好相反)——这样 O 和 m, l 状态在寄存器里增量更新,每个 Q_block 的 O 一次性算完。
但反向的 dV 和 dK 是按 K 维度累加的:
如果用"外层 Q 内层 K",每个 K_block 的 dV 会被多次部分更新——必须用 atomic 或多次 kernel launch。
反向应该用"外层 K,内层 Q":
for k_idx in range(0, N, Bc):
K_block = K[k_idx : k_idx + Bc]
V_block = V[k_idx : k_idx + Bc]
dV_acc = zeros(Bc, d)
dK_acc = zeros(Bc, d)
for q_idx in range(0, N, Br):
Q_block = Q[q_idx : q_idx + Br]
dO_block = dO[q_idx : q_idx + Br]
LSE_block = LSE[q_idx : q_idx + Br] # 来自前向
# 1) 重计算 S = Q @ K^T
S = Q_block @ K_block.T
# 2) 重计算 P (用前向存的 LSE)
P = exp(S - LSE_block.unsqueeze(1))
# 3) D = rowsum(O ⊙ dO) -- 用前向的 O
D_block = sum_per_row(O[q_idx:q_idx+Br] * dO_block)
# 4) dV 累加
dV_acc += P.T @ dO_block
# 5) dP, dS
dP = dO_block @ V_block.T
dS = P * (dP - D_block.unsqueeze(1))
# 6) dQ -- 这一步需要 atomic 累加 (跨 k_idx 多次更新)
dQ_partial = dS @ K_block
atomic_add(dQ[q_idx : q_idx + Br], dQ_partial)
# 7) dK 累加
dK_acc += dS.T @ Q_block
dV[k_idx : k_idx + Bc] = dV_acc
dK[k_idx : k_idx + Bc] = dK_acc
这种循环方式的特点:
- dV 和 dK 对固定的 K_block 在内层 Q 循环中累加,每个 K_block 处理完一次性写出,不需要 atomic。
- dQ 的部分和在内层 Q 循环中逐块算出,但不同 k_idx 的 dQ 部分会落在同一个 q_idx 上——需要 atomic add 或者用第二个 kernel 汇总。
dQ 的处理是 FA2 反向的一个工程难点。
16.3 dQ 的处理:atomic vs 二阶段
dQ 累加的两种工程方案:
16.3.1 方案 A:atomic add 直接写
atomicAdd(dQ + q_idx * d + col, dQ_partial);
简单但慢——全局内存的 atomic 要走到 L2 的原子单元去执行,每个 (K 块, Q 块) 对要提交 个,全程约 个( 时就是 N×N 量级),加起来很多。
实际工业实现用 atomic 时会做几个优化:
- Atomic 在 SMEM 上,最后一次性写 HBM:每个 k_block 的 q 部分只写一次 SMEM atomic,最后写 HBM。但这要求所有 k_block 都在同一个 thread block 内——不现实,因为外层循环遍历所有 K。
- 半精度 atomic:HBM 上半精度 atomic 的字节数少一半。硬件支持并不新——
atomicAdd(__half2*)从 Pascal(sm_60)起就有,atomicAdd(__half*)从 Volta(sm_70)起就有——真正的障碍是精度:dQ 是要跨几十上百个 k_block 累加的量,用 fp16 累加会明显掉精度。
16.3.2 方案 B:两阶段 kernel
第一个 kernel 算 dV 和 dK(外层 K 循环结构),把 dS 写到一个临时缓冲。
第二个 kernel 把 dQ = dS @ K 算出来(不同 K_block 的 dS 已经分别 stage)。
但这要求中间存储 dS——而 dS 又是 N×N 的 attention 矩阵,就是 FA 想避免的东西!
FA2 官方实现走的是方案 A 的一个变体——atomicAdd 到一块 fp32 的 dQ 累加缓冲,而不是直接 atomic 到最终的 dQ。flash-attn/csrc/flash_attn/src/flash_bwd_kernel.h:678 就是这一行:
for (int i = 0; i < size(acc_dq); ++i) { atomicAdd(&tdQgdQaccum(i), acc_dq(i)); }
tdQgdQaccum 指向 params.dq_accum_ptr,元素类型是 ElementAccum(fp32);主循环结束后再由一个单独的 kernel(flash-attn/csrc/flash_attn/src/flash_bwd_preprocess_kernel.h:185 的 convert_dQ)乘上 softmax scale、转成输出 dtype 写进 dQ。这样做的理由:
- 精度:累加在 fp32 上做,跨多少个 k_block 都不掉精度。
- 不等返回值:源码是逐元素对 fragment 里的每个 fp32 调
atomicAdd,但返回值没人用,编译器会生成不等结果的归约指令(nvcc 13.4 编一个同样写法的小 kernel,sm_80 上每个元素一条RED.E.ADD.F32,sm_90a 上是REDG.E.ADD.F32;不同 CUDA 版本可能不同),线程发出去就能接着算下一块。 - 确定性可选:
params.deterministic打开时,每个 thread block 会 atomic 到各自独立的一块 dQ_accum——flash-attn/csrc/flash_attn/src/flash_bwd_kernel.h:125那一行+ (!params.deterministic ? 0 : blockIdx.x * params.dq_accum_split_stride)就是全部机关(最后由 convert kernel 按固定顺序把各块加起来),用额外显存换掉浮点加法顺序不定带来的不可重现性。Python 接口里deterministic=False,默认关闭。
比中间矩阵方案省下的 HBM 流量,远大于这些 atomic 的开销——这是这个设计能成立的根本原因。
16.4 LSE 在反向中的角色
注意 16.2 节的伪代码中第 2 步 P = exp(S - LSE_block.unsqueeze(1))。
这里 LSE 是前向输出的副产物:。
为什么需要 LSE?因为反向重计算 P 需要前向算过的归一化常数。如果不存 LSE,反向时要重新 online softmax 一遍——多一次完整的 K 维度遍历。
存 LSE 的代价:每个 (batch, head) N 个 fp32,N=4096 时 16 KB,可以忽略。
16.5 反向的 SMEM 与寄存器布局
反向比前向需要更多状态:
- Q_block, K_block, V_block:和前向一样
- dO_block:新增 (Br, d) 大小
- LSE_block:新增 (Br) 大小
- O_block:新增 (Br, d)(用来算 D = rowsum(O ⊙ dO))
- dV_acc, dK_acc:(Bc, d) 累加器,保留在寄存器
SMEM 占用更紧,dV_acc/dK_acc 在寄存器中。寄存器压力比前向更大。
16.6 反向 Kernel 骨架
template <int Br, int Bc, int d>
__global__ void flash_attn_bwd(
const half* Q, const half* K, const half* V, const half* O, const half* dO,
const float* LSE,
float* dQ_accum, half* dK, half* dV, // dQ 先累加到 fp32 缓冲,见 16.3 节
int N, float softmax_scale, bool is_causal
) {
// Grid: (ceil(N / Bc), H, B);以下指针运算省略了 batch/head 偏移和 N 不整除时的越界处理
int k_tile_idx = blockIdx.x;
int head_idx = blockIdx.y;
int batch_idx = blockIdx.z;
extern __shared__ half smem[];
half* sK = smem;
half* sV = sK + Bc * d;
half* sQ = sV + Bc * d;
half* sO = sQ + Br * d;
half* sdO = sO + Br * d;
float* sLSE = (float*)(sdO + Br * d);
// Load K, V tile (一次, 整个 inner 循环用)
load_tile(sK, K + k_tile_idx * Bc * d, Bc * d);
load_tile(sV, V + k_tile_idx * Bc * d, Bc * d);
__syncthreads();
// 寄存器中的 dV, dK 累加器
float dV_acc[MMAS_K][MMAS_D][4] = {0};
float dK_acc[MMAS_K][MMAS_D][4] = {0};
// 内层: 遍历所有 Q tile
int q_start = is_causal ? k_tile_idx * Bc : 0;
for (int q_tile = q_start; q_tile < N; q_tile += Br) {
// Load Q, O, dO, LSE
load_tile(sQ, Q + q_tile * d, Br * d);
load_tile(sO, O + q_tile * d, Br * d);
load_tile(sdO, dO + q_tile * d, Br * d);
load_lse(sLSE, LSE + q_tile, Br);
__syncthreads();
// 1) S = Q @ K^T (与前向一样)
float S_frag[MMAS_M][MMAS_N][4];
compute_QKt(S_frag, sQ, sK);
scale_and_mask(S_frag, softmax_scale, q_tile, k_tile_idx, is_causal);
// 2) P = exp(S - LSE) -- 比 online softmax 简单, 因为 LSE 已知
float P_frag[MMAS_M][MMAS_N][4];
for (int i = 0; i < MMAS_M; ++i) {
for (int j = 0; j < MMAS_N; ++j) {
for (int e = 0; e < 4; ++e) {
int row = ...; // lane 持有的行
P_frag[i][j][e] = expf(S_frag[i][j][e] - sLSE[row]);
}
}
}
// 3) D = rowsum(O ⊙ dO)
float D_local[MMAS_M][2] = {0};
compute_D(D_local, sO, sdO); // load O, dO 的 fragment 并按行 reduce
// 4) dV_acc += P^T @ dO
// 注意 P 是 fp16, dO 是 fp16, dV_acc 是 fp32
unsigned P_fp16[MMAS_M][MMAS_N];
cast_fp32_to_fp16(P_fp16, P_frag);
unsigned dO_frag[MMAS_M][MMAS_D];
load_dO_fragment(dO_frag, sdO);
// 注意是 P^T, 需要 ldmatrix.trans 或者交换 fragment 角色
accumulate_matmul_T(dV_acc, P_fp16, dO_frag);
// 5) dP = dO @ V^T
float dP_frag[MMAS_M][MMAS_N][4];
compute_dOVt(dP_frag, sdO, sV);
// 6) dS = P ⊙ (dP - D)
for (int i = 0; i < MMAS_M; ++i) {
for (int j = 0; j < MMAS_N; ++j) {
for (int e = 0; e < 4; ++e) {
int row_local = e / 2;
dP_frag[i][j][e]
= P_frag[i][j][e] * (dP_frag[i][j][e] - D_local[i][row_local]);
}
}
}
// 现在 dP_frag 实际是 dS_frag
// 7) dQ_partial = dS @ K, atomic add to dQ
unsigned dS_fp16[MMAS_M][MMAS_N];
cast_fp32_to_fp16(dS_fp16, dP_frag);
unsigned K_frag[MMAS_N][MMAS_D];
load_K_fragment(K_frag, sK);
float dQ_partial[MMAS_M][MMAS_D][4] = {0};
accumulate_matmul(dQ_partial, dS_fp16, K_frag);
// atomic add 到 fp32 的 dQ_accum;乘 softmax_scale、转 half 由后续 kernel 做
atomic_write_dQ(dQ_accum + q_tile * d, dQ_partial);
// 8) dK_acc += dS^T @ Q
unsigned Q_frag[MMAS_M][MMAS_D];
load_Q_fragment(Q_frag, sQ);
accumulate_matmul_T(dK_acc, dS_fp16, Q_frag);
__syncthreads();
}
// 写 dV, dK(dK 写出前要乘 softmax_scale,因为 S = softmax_scale · QK^T)
write_dV(dV + k_tile_idx * Bc * d, dV_acc);
write_dK(dK + k_tile_idx * Bc * d, dK_acc);
}
注意:上面是示意骨架,
load_tile、compute_QKt等辅助函数和MMAS_*常量都没有给出,row的计算也留空,不能直接编译。完整实现请参考flash-attn/csrc/flash_attn/src/flash_bwd_kernel.h(GitHub 上的同一份)。
16.7 反向的性能与挑战
反向比前向慢的原因:
- 五个 GEMM vs 前向的两个:重算 、,再加 、、 三个梯度各一个,合计 FLOPs——正好是前向 的 2.5 倍。光这一条就解释了大半的耗时比。
- dQ 的 atomic write:每次 inner Q 循环一次 atomic,N 大时累计严重。
- 更多 SMEM 输入:Q, K, V, O, dO, LSE 全部要在 SMEM 中。
- 重计算开销:S 和 P 在反向时要重算(这是 FA 的设计选择,避免存中间矩阵)。
FLOPs 比是 2.5,而反向在单位算力上还跑得比前向更不满:FA2 论文在 A100 上测得前向最高约 73% 峰值、反向最高约 63%(论文正是按前向 × head 数、反向再乘 2.5 来折算 TFLOPs 的),dQ 的 atomic 和更大的 SMEM 需求就是反向跑不满的原因。两者合起来,按两者的最高利用率粗估,反向耗时约为前向的 倍,具体比值随序列长度、head dim、是否 causal 而变。本专栏没有条件实测,读者自己跑一遍 flash_attn_func 的 forward / backward 计时比就能看到这个比值。
16.8 工业级实现的关键技巧
FlashAttention 官方实现(Dao-AILab/flash-attention)做了大量工程优化。要区分两支代码:
- FA2(
flash-attn/csrc/flash_attn/src/,sm80 起):同步用的还是__syncthreads()——grep 整个flash-attn/csrc/flash_attn/src/flash_bwd_kernel.h找不到一处mbarrier(0 处)。 - FA3(
hopper/,sm90 专属):才用上 mbarrier / 命名 barrier 驱动的 producer-consumer 流水(第 17 章)。
FA2 反向这一支真正的技巧集中在别处:
- 前置的 preprocess kernel:
flash-attn/csrc/flash_attn/src/flash_bwd_preprocess_kernel.h:58的compute_dot_do_o在主循环之前把 一次算完存好(逐行点积在同文件 :25 的dot_do_o),主循环直接读,省掉每轮重算。 - causal 时裁掉整块:q < k 的 tile 整块跳过,长序列上省掉接近一半工作。
- dQ 累加缓冲用 fp32:见 16.3 节,精度与性能的折中点在这里。
- 确定性开关:
deterministic用额外显存换可重现的梯度。 - SMEM 的精打细算:反向的 SMEM 压力比前向大得多,源码里有多处刻意的复用与
__syncthreads()保护:V 放进寄存器(Is_V_in_regs)时 sdS 直接复用 sV 的位置(flash-attn/csrc/flash_attn/src/flash_bwd_kernel.h:169),sP 与 sdQ 共用同一块内存(:176 的注释 "sP and sdQ share the same memory so be careful"),下一块 dO 要预取进刚被 dV 的 GEMM 读过的同一个 sdO,所以先同步(:640 的注释 "Need syncthreads since we're writing to the same sdO location")。
每一个都是精度、显存、性能之间的权衡。
16.9 这一章的小结与下一章
FA2 反向的核心要点:
- 数学:dV/dK/dQ 都涉及重计算 S 和 P——这是用计算换存储的核心 trade-off。
- 循环顺序翻转:外层 K,内层 Q。dV/dK 在外层累加,dQ 必须 atomic。
- LSE 是前向到反向的桥梁:让反向不需要重新 online softmax。
- D = rowsum(O ⊙ dO) 是简化 dS 公式的关键。
- 反向的 FLOPs 是前向的 2.5 倍:五个 GEMM(,前向 );再加上 dQ 的 atomic 累加和更大的 SMEM 压力让利用率更低,按 FA2 论文的利用率粗估,耗时约为前向的 3 倍。
到第 16 章为止,读者已经能写出一个能用、性能合理的 FA2(forward + backward)。但还没用到 Hopper 的杀手锏——TMA 和 Warp Specialization。
第 17 章我们把 FA2 重写到 Hopper 的最优形态——TMA 异步拷贝代替 cp.async、WGMMA 代替 mma.sync、Producer/Consumer warp 流水线。读完第 17 章读者会理解为什么 FA3 在 H100 上能跑到 ~740 TFLOPs(FP16,75% 利用率),以及"现代 GPU kernel"和"Ampere 时代的 GPU kernel"在编程模型上的根本差异。
本章动手练习:
- 推导 dS = P ⊙ (dP - D) 的代数过程。验证 D 的两种等价定义。
- 实现 16.6 节骨架,测反向 / 前向的耗时比,看它比 2.5 倍的 FLOPs 比高出多少。
- 思考:如果不存 LSE,反向时怎么重新算?需要多少额外 HBM 流量?