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

第 14 章 Attention 的访存瓶颈与 IO-Aware 思想

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

FlashAttention 改的不是 attention 的数学,是它的 I/O。 这一章先把朴素实现的 HBM 流量一笔笔算清楚——IO-aware 的必要性才不是一句口号。

14.1 Attention 的数学

经典 scaled dot-product attention:

Attention(Q,K,V)=softmax(QKTd)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d}}\right) V

形状:

  • Q∈RN×dQ \in \mathbb{R}^{N \times d}(N 个 query,每个 d 维)
  • K∈RN×dK \in \mathbb{R}^{N \times d}(N 个 key,每个 d 维)
  • V∈RN×dV \in \mathbb{R}^{N \times d}(N 个 value,每个 d 维)
  • 输出 O∈RN×dO \in \mathbb{R}^{N \times d}

中间结果:

  • S=QKT∈RN×NS = QK^T \in \mathbb{R}^{N \times N}(attention scores)
  • P=softmax(S)∈RN×NP = \text{softmax}(S) \in \mathbb{R}^{N \times N}(attention probabilities)
  • O=PVO = PV

数学上每个步骤都很清晰。问题在于在 GPU 上这种逐步计算的 HBM 流量爆炸。

14.2 标准 Attention 的 HBM 流量精确计算

设 N=4096N = 4096(序列长度),d=64d = 64(head dim),FP16 数据。本章的 KB、MB 按 2102^{10}、2202^{20} 字节计。

Step 1: S=QKTS = QK^T

  • 读 Q:N×d×2=512N \times d \times 2 = 512 KB
  • 读 K:N×d×2=512N \times d \times 2 = 512 KB
  • 写 S:N×N×2=32N \times N \times 2 = 32 MB

Step 2: P=softmax(S)P = \text{softmax}(S)

朴素 softmax 3 pass:

  • 读 S:3 次 × 32 MB = 96 MB
  • 写 P:32 MB

Step 3: O=PVO = PV

  • 读 P:32 MB
  • 读 V:512 KB
  • 写 O:512 KB

总 HBM 流量:

0.5 + 0.5 + 32 (写 S)
+ 96 + 32 (softmax)
+ 32 + 0.5 + 0.5 (写 O)
= 194 MB

而计算量只有:

  • QKTQK^T: 2N2d=2×40962×64≈2.15⋅1092 N^2 d = 2 \times 4096^2 \times 64 \approx 2.15 \cdot 10^9 FLOPs
  • PVPV: 2N2d≈2.15⋅1092 N^2 d \approx 2.15 \cdot 10^9 FLOPs
  • 合计:4.3⋅1094.3 \cdot 10^9 FLOPs = 4.3 GFLOPs

实际算术强度 = 4.3 GFLOPs / 194 MB(2.03⋅1082.03 \cdot 10^8 字节)≈ 21 FLOPs/byte

把这个数字放到 H100 SXM 的 Roofline 上:

  • 21 < 295(BF16/FP16 Tensor Core 稠密口径的临界点,989 TFLOPs ÷ 3.35 TB/s)
  • 所以 attention 是带宽 bound

理论上限 = 21 × 3.35 TB/s ≈ 71 TFLOPs。但 Tensor Core 稠密峰值 989 TFLOPs——理论上 attention 只能用到约 7% 算力。

本专栏没有 GPU 可实测,这里只给结论的方向:分步实现连这 71 TFLOPs 的上限都够不着——实际带宽利用率到不了 100%,中间还要加上多次 kernel launch 的开销。即使把 softmax 融合成只读一遍 S,总流量也还有 130 MB,AI ≈ 32(即第 4 章按 S、P 各写读一次估的 d/2d/2),上限约 106 TFLOPs,约合峰值的 11%。朴素 attention 的算力利用率只有百分之几到一成左右,这一点与是不是「写得够仔细」无关,是访存结构决定的。

14.3 中间矩阵 S 和 P 的代价

仔细看 14.2 节的流量分解,会发现一个惊人的事实:194 MB HBM 流量中,绝大部分是 S 和 P 矩阵的反复读写:

真正"内容"流量:
  Q: 0.5 MB    K: 0.5 MB    V: 0.5 MB    O: 0.5 MB
  合计: 2 MB

中间矩阵流量:
  写 S: 32 MB
  softmax 读 S 3次 + 写 P: 96+32 = 128 MB (融合成一遍的实现只读一次 S, 这一项降到 64 MB)
  PV 读 P: 32 MB
  合计: 192 MB

中间流量 / 总流量 = 192 / 194 ≈ 99%

约 99% 的 HBM 带宽消耗在中间矩阵上——而这些矩阵的存在仅仅是因为我们一步一步串行算。如果能把它们完全消除,HBM 流量的理想下界就只剩 Q/K/V/O 各读写一次的 ~2 MB。注意这个下界要求片上存储放得下整份 K/V(FA1 论文里 M=Θ(Nd)M = \Theta(Nd) 的那一端);SRAM 容量有限时,论文 Proposition 3 证明了不存在对所有 M∈[d,Nd]M \in [d, Nd] 都只需 o(N2d2/M)o(N^2 d^2 / M) 次 HBM 访问的精确 attention 算法。14.6 节会看到 FA 因为 K/V 要被反复重读,实际落点远高于 2 MB。

这就是 FlashAttention 的核心 insight。

14.4 IO-Aware:FA 的核心思想

FlashAttention 论文(Dao et al., 2022)的关键 idea:不要把 S 和 P 写到 HBM。

具体做法:

  1. Tile 化 attention:把 Q,K,VQ, K, V 沿 N 维切成 tile(block)。
  2. 每次只对一组 (Q_tile, K_tile, V_tile) 计算:S_tile 在片上生成(论文写作 on-chip SRAM;FA2 的实现里 S_tile 就留在寄存器),softmax_tile 在片上完成,PV 累加到 O_tile,完全不写中间矩阵。
  3. 跨 K_tile 的 softmax 合并:用第 6 章讲的 online softmax,把多个 K_tile 的结果在线合并。
flowchart LR
  subgraph Naive [朴素 Attention · 多 kernel]
    N1[QK^T → S in HBM]
    N2[softmax S → P in HBM]
    N3[PV → O in HBM]
    N1 --> N2 --> N3
  end
  subgraph FA [FlashAttention · 单 kernel]
    F1[Tile by tile]
    F1 --> F2[S_tile on-chip]
    F2 --> F3[Online softmax on-chip]
    F3 --> F4[Accumulate to O fragment in register]
    F4 --> F5[全部 K_tile 处理完后, 写 O 一次]
  end

14.5 FA 的算法骨架(Forward)

先把一件常被写错的事说清楚:「外层 Q、内层 K/V」是 FlashAttention v2 的循环顺序,不是 FA1 的。

  • FA1(Dao et al., 2022)论文 Algorithm 1:外层遍历 K/V 块,内层遍历 Q 块。K/V 块进 SMEM 后被所有 Q 块复用,代价是 Q、O、以及 m,ℓm,\ell 统计量每轮都要重新读回来、再写回去。
  • FA2(Dao, 2023)论文 3.2 节明写这是 swapping the order of the loop(outer loop over row blocks and inner loop over column blocks):外层 Q、内层 K/V。Q 只读一次,O 和 m,ℓm,\ell 全程留在寄存器里,只在最后写一次 HBM。

下面这段伪代码是 FA2 的形态——也就是本专栏第 15 章要手写的那一版:

# FA2 Forward 算法 (Dao, 2023, FlashAttention-2 论文 Algorithm 1)
# 外层 Q、内层 K/V —— 注意 FA1 的循环顺序与此相反
# 输入: Q, K, V ∈ [N, d]
# 输出: O ∈ [N, d], LSE ∈ [N]

# 沿 N 维分块
B_q = 64   # Query block size
B_k = 64   # Key block size

# 外层: 遍历 Q 的 block
for q_idx in range(0, N, B_q):
    Q_block = Q[q_idx : q_idx + B_q]          # [B_q, d]

    # 累加状态
    O_block = zeros(B_q, d)
    m_block = -inf * ones(B_q)                 # 行级 max
    l_block = zeros(B_q)                       # 行级 sum

    # 内层: 遍历 K 的 block
    for k_idx in range(0, N, B_k):
        K_block = K[k_idx : k_idx + B_k]      # [B_k, d]
        V_block = V[k_idx : k_idx + B_k]      # [B_k, d]

        # 1) 计算 S_block = Q_block @ K_block^T
        S_block = Q_block @ K_block.T          # [B_q, B_k]

        # 2) Online softmax 更新 m, l
        #    (本伪代码不带 mask; 若加 mask 可能使某行 m_new 仍为 -inf,
        #     须按第 6 章做 -INF 防护, 否则 exp(-inf - (-inf)) = NaN)
        m_new = max(m_block, max_per_row(S_block))   # [B_q]
        P_block = exp(S_block - m_new)               # [B_q, B_k]
        alpha = exp(m_block - m_new)                  # [B_q]
        l_new = alpha * l_block + sum_per_row(P_block)

        # 3) 累加输出
        O_block = alpha * O_block + P_block @ V_block

        m_block = m_new
        l_block = l_new

    # 最终归一化
    O[q_idx : q_idx + B_q] = O_block / l_block
    LSE[q_idx : q_idx + B_q] = m_block + log(l_block)

关键点:

  1. 没有 N×N 矩阵:S 和 P 只在内层循环的当前 K_tile 内存在,B_q × B_k 大小(很小)。这一条 FA1 和 FA2 共有,是 IO-aware 的本体。
  2. O 在内层循环中累加、外层循环结束才写出:每个 K_tile 都对当前的 O_block 贡献一份,用 online softmax 的 alpha 修正之前的累积。因为外层是 Q,O_block 从头到尾待在同一组寄存器里——这正是 FA2 换循环顺序换来的好处。
  3. 最后才做归一化:除以最终的 l_block。
  4. LSE(log-sum-exp)作为副产物输出:反向传播时需要。

14.6 FA 的 HBM 流量重新计算

把 FA 的 HBM 流量算一遍。N=4096, d=64, FP16:

  • 读 Q:每个 q_block 读 1 次,总共读 1 次 = N×d×2N \times d \times 2 = 512 KB
  • 读 K:对每个 q_block,遍历所有 k_block,总共读 N/BqN/B_q 次 K = N×N/Bq×d×2N \times N/B_q \times d \times 2
  • 读 V:同 K,读 N/BqN/B_q 次

N/Bq=64N/B_q = 64,所以 K 读 64 次 = 32 MB,V 读 64 次 = 32 MB。

  • 写 O:1 次 = 512 KB
  • 写 LSE:很小(N 个 fp32)

总 HBM 流量 ≈ 0.5 + 32 + 32 + 0.5 ≈ 65 MB

比朴素的 194 MB 少了三分之二。省掉的是什么?是中间矩阵(S, P)的 HBM 读写——这些是带宽里"真正不应该有"的部分。

但同时 FA 引入了一个新代价:K 和 V 被多次读(每个 q_block 都读一次 K 和 V)。这个代价在 Bq=64B_q = 64、N=4096N = 4096 时是 64×。

在这一组具体参数下,总 HBM 流量从 194 MB 降到 65 MB。实际算术强度:

  • FLOPs 不变:4.3 GFLOPs
  • HBM:65 MB(6.8⋅1076.8 \cdot 10^7 字节)
  • AI = 4.3 GFLOPs / 65 MB ≈ 63 FLOPs/byte

仍然 < 295(带宽 bound)。这一组参数下的理论上限:63 × 3.35 TB/s ≈ 211 TFLOPs。这个上限是按「每次 K/V 重读都落到 HBM」算的,偏保守:一个 head 的 K、V 合计才 1 MB,实际重读有相当一部分能命中 H100 SXM5 的 50 MB L2,块也可以取得比 64 更大。

别把这个算例当成普适结论。 上面的 65 MB 是 Bq=64B_q = 64 这一个取值算出来的。FA 的 HBM 访存量的正确写法是 Θ(N2d2/M)\Theta(N^2 d^2 / M)(FA1 论文 Theorem 2,MM 是 SRAM 容量,块大小 Bq,BkB_q, B_k 由 MM 决定),朴素实现是 Θ(Nd+N2)\Theta(Nd + N^2)——两者都随 N2N^2 增长,FA 省下的是一个 M/d2M/d^2 量级的常数因子。按 H100 每 SM 228 KB SMEM(约 11.7 万个 FP16 元素)粗估,d=64d = 64 时约 28,d=128d = 128 时约 7,即十倍量级;FA1 论文在 A100 上实测 GPT-2 medium 前向加反向的 HBM 读写从 40.3 GB 降到 4.4 GB(Figure 2),约 9 倍。要区分的另一个量是显存占用:那个确实从 O(N2)O(N^2) 降到了 O(N)O(N)(N×N 矩阵根本不落地,只留 mm、ℓ\ell)。这两个量常被混为一谈,《Transformer 解剖:从 Attention 到推理系统》专栏第 18 章对它们做了同样的区分。

FA1 / FA2 / FA3 在同一块卡上的差距不来自"算的东西变了"——三者的输出在数学上等价——而来自并行划分和硬件特性的利用:FA2 换了循环顺序、削掉了非矩阵乘 op、改了 warp 间的工作划分(第 15 章);FA3 换用 Hopper 的 TMA + WGMMA + warp specialization(第 17 章)。具体数字见第 17 章 17.8 节,那里给出了带论文出处的口径。

14.7 K/V 重读的代价:能不能进一步优化

FA2 Forward 的 K/V 重读看似浪费。能不能把 K/V 也只读一次?

答案是可以,把外层循环换回 外层遍历 K/V,内层遍历 Q——这正是 FA1 论文 Algorithm 1 的原始写法。代价是 Q 被多次读,且 O、mm、ℓ\ell 每轮都要读回来再写回去(多个 k_block 会先后贡献到同一批 q 行)。FA2 把两层循环对调:外层的各个 Q 块互不依赖,可以直接分给不同 thread block 并行,不需要块间通信;论文 3.2 节给出的动机是在 batch × head 较小时,靠沿序列长度的并行提高 occupancy。O、mm、ℓ\ell 不再反复读写 HBM 是顺带的收益。

两种循环方式各有归宿:

外层循环 重读的是谁 谁在用
外 K/V 内 Q Q、O、mm/ℓ\ell 反复读写 FA1 前向(论文 Algorithm 1)
外 Q 内 K/V K、V 重读 FA2/FA3 前向
外 K/V 内 Q Q 重读,dQ 需 atomic 累加 FA2 反向(第 16 章)

注意最后一行:FA2 前向把循环顺序换成了外 Q,反向却又换了回去——原因在第 16 章。

14.7.1 长序列的进一步优化

如果 N 极大(比如 N=64K,长上下文),而 batch × head 又不足以填满全部 SM,就该把 K/V 维度也切开给多个 block 同时算,每个 block 出一份 partial O 和 partial LSE,最后跨 block 合并——这就是 Split-KV / Flash-Decoding。

它不是 FA3 才有的东西:FlashAttention 仓库在 FA2(sm80)这一支就有专门的一组 split kernel(flash-attn/csrc/flash_attn/src/flash_fwd_split_hdim64_fp16_sm80.cu 等,按 headdim × dtype × causal × 对齐情况铺开了几十个编译单元),Python 侧由 flash_attn_with_kvcache(..., num_splits=...) 暴露(flash-attn/flash_attn/flash_attn_interface.py:1485,参数在 :1503)。文档里写得很清楚:num_splits == 1 不切,> 1 按这个数切,== 0 走启发式自动决定(flash-attn/flash_attn/flash_attn_interface.py:1581)——所以默认值 0 并不是"关掉",而是"交给启发式"。这条路径主要服务 long-context decoding——输入 128K、每次只生成一个 token,每个 head 的 Q 只有一行(GQA 下一个 KV head 也只对应组内那几个 query head,第 4 章算过此时算术强度约等于组数),不切 K/V 就只有极少数 SM 在干活。第 18 章会把它展开。

14.8 关于 IO-Aware 的更广义理解

FlashAttention 的成功不只是一个算法。它代表了一种思维方式:

重写算法的数据流,而不是重写算法本身。

数学上,FA 算的是同一个 attention(输出在浮点误差范围内相等)。但它的"数据流"完全重组了——把 N×N 中间矩阵从 HBM 挤出去,让 K/V 多读几次换中间矩阵不写。这种"用一种带宽换另一种带宽"的 trade-off 是 GPU 算法设计的核心模式。

围绕内存与带宽重新组织计算的思路,在其他地方也能看到:

  • FlashConv(H3 论文):把状态空间模型里的 FFT 长卷积融合进片上 SRAM,同样是 IO-aware 的写法
  • PagedAttention:按页管理 KV cache,解决的是显存容量与碎片问题,不是层级间搬运,但同样是「按硬件内存特性组织数据」
  • Speculative Decoding:decode 是带宽 bound 的,一次读权重同时验证多个候选 token,把一份 HBM 流量摊给更多计算

它们各自针对的瓶颈不同,共同点是先看清数据在内存里怎么流动,再决定计算怎么组织。读懂 FA 之后,读者会发现这种思维在 LLM 系统的每一层都有应用。

14.9 这一章的小结与下一章

这一章把 attention 的访存瓶颈彻底剖开:

  1. 朴素 attention 是带宽 bound(算例 AI≈21):HBM 流量主要消耗在 N×N 中间矩阵 S 和 P 上。
  2. FA 的 idea 是把中间矩阵留在片上:用 online softmax 跨 K-tile 累积。显存占用因此从 O(N2)O(N^2) 降到 O(N)O(N)。
  3. HBM 访存量则是 Θ(N2d2/M)\Theta(N^2 d^2/M):仍随 N2N^2 增长,省下的是 M/d2M/d^2 量级的常数因子(十倍量级),不是某个固定比例。14.6 节那个 194 MB → 65 MB 只是 Bq=64B_q=64 的一个算例。
  4. 算例参数下的理论上限是 ~211 TFLOPs(仍带宽 bound):这是按 Bq=64B_q=64、K/V 重读全落 HBM 估的保守值;FA3 用更大的块和 Hopper 特性,实测吞吐在这个量级之上(第 17 章)。
  5. IO-Aware 是一种思维方式:在 LLM 系统的每一层都有应用。

第 15 章我们正式动手——把 14.5 节的伪代码翻译成具体的 CUDA kernel。我们会用第三篇 GEMM 优化的所有工具(Tensor Core、ldmatrix、SMEM tile、double buffer)来组装 FA2 前向。读完第 15 章读者会拥有一个结构完整的 FA2 forward kernel 骨架;它的性能水位见 15.6 节:官方 FA2 在 H100 上约 35% Tensor Core 利用率(FA3 论文口径),手写骨架只会更低。

本章动手练习:

  1. 用 PyTorch 写朴素 attention 和调用 torch.nn.functional.scaled_dot_product_attention(CUDA 上按条件在 FlashAttention、cuDNN、memory-efficient 等后端之间选择,可用 torch.nn.attention.sdpa_kernel(SDPBackend.FLASH_ATTENTION) 指定走 FA),用 Nsight Compute 测两者的 HBM 读写流量,验证差距。
  2. 推导 FA 反向的 HBM 流量公式,对比朴素反向(需要 N×N 中间梯度矩阵)。
  3. 把 FA1 论文 Algorithm 1(外 K/V 内 Q)和 FA2 论文 Algorithm 1(外 Q 内 K/V)并排放,逐行找出两者的差异——尤其是 mm、ℓ\ell、O 分别在哪一层循环里活着。