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

第 16 章 FA2 反向:dQ/dK/dV 的重计算

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

反向不是把前向倒着写一遍:前向用着顺手的数据布局,到反向几乎每一处都要重新安排。 这一章讲 dQ/dK/dV 三条梯度各自的并行策略,以及为什么 S/P 必须重算而不是存下来。

16.1 反向的数学:从 dO 到 dQ/dK/dV

设 attention 前向:

S=QKT,P=softmax(S),O=PVS = QK^T, \quad P = \text{softmax}(S), \quad O = PV

给定 dO=∂L/∂OdO = \partial L / \partial O,我们要算:

dQ=∂L/∂Q,dK=∂L/∂K,dV=∂L/∂VdQ = \partial L / \partial Q, \quad dK = \partial L / \partial K, \quad dV = \partial L / \partial V

链式法则一步步:

dV 最直接:

dV=PT⋅dOdV = P^T \cdot dO

dP:

dP=dO⋅VTdP = dO \cdot V^T

dS(softmax 的反向): 设 pi=softmax(si)p_i = \text{softmax}(s_i),则 ∂pi∂sj=pi(δij−pj)\frac{\partial p_i}{\partial s_j} = p_i (\delta_{ij} - p_j)。所以:

dSij=Pij⋅(dPij−∑kPikdPik)dS_{ij} = P_{ij} \cdot (dP_{ij} - \sum_k P_{ik} dP_{ik})

定义 Di=∑kPikdPik=∑kPik(dOVT)ikD_i = \sum_k P_{ik} dP_{ik} = \sum_k P_{ik} (dO V^T)_{ik},则:

dSij=Pij⋅(dPij−Di)dS_{ij} = P_{ij} \cdot (dP_{ij} - D_i)

dQ, dK:

dQ=dS⋅K,dK=dST⋅QdQ = dS \cdot K, \quad dK = dS^T \cdot Q

总结,反向需要算的量:

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

注意 DD 这一项——它是 P⊙dPP \odot dP 按行求和。论文里有个简化技巧:

Di=∑kPik dPik=∑kPik (dOi⋅Vk)=dOi⋅(∑kPikVk)=dOi⋅OiD_i = \sum_k P_{ik} \, dP_{ik} = \sum_k P_{ik} \, (dO_i \cdot V_k) = dO_i \cdot \Big(\sum_k P_{ik} V_k\Big) = dO_i \cdot O_i

注意这里的 PP 是已归一化的 softmax 输出,Oi=∑kPikVkO_i = \sum_k P_{ik} V_k 正是前向的输出行——所以这是恒等式,不带任何额外系数。

更简洁:D=rowsum(O⊙dO)D = \text{rowsum}(O \odot dO)。这避免了显式算 P,只需要前向已经写出去的 O 和上游传下来的 dO。

16.2 反向的循环顺序:外层 K,内层 Q

FA2 前向用的是"外层 Q,内层 K/V"(第 14 章 14.5 节;FA1 的前向恰好相反)——这样 O 和 m, l 状态在寄存器里增量更新,每个 Q_block 的 O 一次性算完。

但反向的 dV 和 dK 是按 K 维度累加的:

dVj=∑iPij⋅dOidV_j = \sum_i P_{ij} \cdot dO_i

如果用"外层 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 块) 对要提交 Br×dB_r \times d 个,全程约 N2d/BcN^2 d / B_c 个(d=Bcd = B_c 时就是 N×N 量级),加起来很多。

实际工业实现用 atomic 时会做几个优化:

  1. Atomic 在 SMEM 上,最后一次性写 HBM:每个 k_block 的 q 部分只写一次 SMEM atomic,最后写 HBM。但这要求所有 k_block 都在同一个 thread block 内——不现实,因为外层循环遍历所有 K。
  2. 半精度 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 是前向输出的副产物:LSEi=mi+log⁡ℓi\text{LSE}_i = m_i + \log \ell_i。

为什么需要 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 反向的性能与挑战

反向比前向慢的原因:

  1. 五个 GEMM vs 前向的两个:重算 S=QKTS = QK^T、dP=dO VTdP = dO\,V^T,再加 dVdV、dKdK、dQdQ 三个梯度各一个,合计 10N2d10N^2d FLOPs——正好是前向 4N2d4N^2d 的 2.5 倍。光这一条就解释了大半的耗时比。
  2. dQ 的 atomic write:每次 inner Q 循环一次 atomic,N 大时累计严重。
  3. 更多 SMEM 输入:Q, K, V, O, dO, LSE 全部要在 SMEM 中。
  4. 重计算开销:S 和 P 在反向时要重算(这是 FA 的设计选择,避免存中间矩阵)。

FLOPs 比是 2.5,而反向在单位算力上还跑得比前向更不满:FA2 论文在 A100 上测得前向最高约 73% 峰值、反向最高约 63%(论文正是按前向 4N2d4N^2d × head 数、反向再乘 2.5 来折算 TFLOPs 的),dQ 的 atomic 和更大的 SMEM 需求就是反向跑不满的原因。两者合起来,按两者的最高利用率粗估,反向耗时约为前向的 2.5×73/63≈2.92.5 \times 73/63 \approx 2.9 倍,具体比值随序列长度、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 反向这一支真正的技巧集中在别处:

  1. 前置的 preprocess kernel:flash-attn/csrc/flash_attn/src/flash_bwd_preprocess_kernel.h:58 的 compute_dot_do_o 在主循环之前把 D=rowsum(O⊙dO)D = \text{rowsum}(O \odot dO) 一次算完存好(逐行点积在同文件 :25 的 dot_do_o),主循环直接读,省掉每轮重算。
  2. causal 时裁掉整块:q < k 的 tile 整块跳过,长序列上省掉接近一半工作。
  3. dQ 累加缓冲用 fp32:见 16.3 节,精度与性能的折中点在这里。
  4. 确定性开关:deterministic 用额外显存换可重现的梯度。
  5. 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 反向的核心要点:

  1. 数学:dV/dK/dQ 都涉及重计算 S 和 P——这是用计算换存储的核心 trade-off。
  2. 循环顺序翻转:外层 K,内层 Q。dV/dK 在外层累加,dQ 必须 atomic。
  3. LSE 是前向到反向的桥梁:让反向不需要重新 online softmax。
  4. D = rowsum(O ⊙ dO) 是简化 dS 公式的关键。
  5. 反向的 FLOPs 是前向的 2.5 倍:五个 GEMM(10N2d10N^2d,前向 4N2d4N^2d);再加上 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"在编程模型上的根本差异。

本章动手练习:

  1. 推导 dS = P ⊙ (dP - D) 的代数过程。验证 D 的两种等价定义。
  2. 实现 16.6 节骨架,测反向 / 前向的耗时比,看它比 2.5 倍的 FLOPs 比高出多少。
  3. 思考:如果不存 LSE,反向时怎么重新算?需要多少额外 HBM 流量?