CUDA 算子工程:手写 FlashAttention v2 之路
第 14 章 Attention 的访存瓶颈与 IO-Aware 思想
FlashAttention 改的不是 attention 的数学,是它的 I/O。 这一章先把朴素实现的 HBM 流量一笔笔算清楚——IO-aware 的必要性才不是一句口号。
14.1 Attention 的数学
经典 scaled dot-product attention:
形状:
- (N 个 query,每个 d 维)
- (N 个 key,每个 d 维)
- (N 个 value,每个 d 维)
- 输出
中间结果:
- (attention scores)
- (attention probabilities)
数学上每个步骤都很清晰。问题在于在 GPU 上这种逐步计算的 HBM 流量爆炸。
14.2 标准 Attention 的 HBM 流量精确计算
设 (序列长度),(head dim),FP16 数据。本章的 KB、MB 按 、 字节计。
Step 1:
- 读 Q: KB
- 读 K: KB
- 写 S: MB
Step 2:
朴素 softmax 3 pass:
- 读 S:3 次 × 32 MB = 96 MB
- 写 P:32 MB
Step 3:
- 读 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
而计算量只有:
- : FLOPs
- : FLOPs
- 合计: FLOPs = 4.3 GFLOPs
实际算术强度 = 4.3 GFLOPs / 194 MB( 字节)≈ 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 各写读一次估的 ),上限约 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 论文里 的那一端);SRAM 容量有限时,论文 Proposition 3 证明了不存在对所有 都只需 次 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。
具体做法:
- Tile 化 attention:把 沿 N 维切成 tile(block)。
- 每次只对一组 (Q_tile, K_tile, V_tile) 计算:S_tile 在片上生成(论文写作 on-chip SRAM;FA2 的实现里 S_tile 就留在寄存器),softmax_tile 在片上完成,PV 累加到 O_tile,完全不写中间矩阵。
- 跨 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、以及 统计量每轮都要重新读回来、再写回去。
- 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 和 全程留在寄存器里,只在最后写一次 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)
关键点:
- 没有 N×N 矩阵:S 和 P 只在内层循环的当前 K_tile 内存在,B_q × B_k 大小(很小)。这一条 FA1 和 FA2 共有,是 IO-aware 的本体。
- O 在内层循环中累加、外层循环结束才写出:每个 K_tile 都对当前的 O_block 贡献一份,用 online softmax 的 alpha 修正之前的累积。因为外层是 Q,O_block 从头到尾待在同一组寄存器里——这正是 FA2 换循环顺序换来的好处。
- 最后才做归一化:除以最终的 l_block。
- LSE(log-sum-exp)作为副产物输出:反向传播时需要。
14.6 FA 的 HBM 流量重新计算
把 FA 的 HBM 流量算一遍。N=4096, d=64, FP16:
- 读 Q:每个 q_block 读 1 次,总共读 1 次 = = 512 KB
- 读 K:对每个 q_block,遍历所有 k_block,总共读 次 K =
- 读 V:同 K,读 次
,所以 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)。这个代价在 、 时是 64×。
在这一组具体参数下,总 HBM 流量从 194 MB 降到 65 MB。实际算术强度:
- FLOPs 不变:4.3 GFLOPs
- HBM:65 MB( 字节)
- 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 是 这一个取值算出来的。FA 的 HBM 访存量的正确写法是 (FA1 论文 Theorem 2, 是 SRAM 容量,块大小 由 决定),朴素实现是 ——两者都随 增长,FA 省下的是一个 量级的常数因子。按 H100 每 SM 228 KB SMEM(约 11.7 万个 FP16 元素)粗估, 时约 28, 时约 7,即十倍量级;FA1 论文在 A100 上实测 GPT-2 medium 前向加反向的 HBM 读写从 40.3 GB 降到 4.4 GB(Figure 2),约 9 倍。要区分的另一个量是显存占用:那个确实从 降到了 (N×N 矩阵根本不落地,只留 、)。这两个量常被混为一谈,《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、、 每轮都要读回来再写回去(多个 k_block 会先后贡献到同一批 q 行)。FA2 把两层循环对调:外层的各个 Q 块互不依赖,可以直接分给不同 thread block 并行,不需要块间通信;论文 3.2 节给出的动机是在 batch × head 较小时,靠沿序列长度的并行提高 occupancy。O、、 不再反复读写 HBM 是顺带的收益。
两种循环方式各有归宿:
| 外层循环 | 重读的是谁 | 谁在用 |
|---|---|---|
| 外 K/V 内 Q | Q、O、/ 反复读写 | 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 的访存瓶颈彻底剖开:
- 朴素 attention 是带宽 bound(算例 AI≈21):HBM 流量主要消耗在 N×N 中间矩阵 S 和 P 上。
- FA 的 idea 是把中间矩阵留在片上:用 online softmax 跨 K-tile 累积。显存占用因此从 降到 。
- HBM 访存量则是 :仍随 增长,省下的是 量级的常数因子(十倍量级),不是某个固定比例。14.6 节那个 194 MB → 65 MB 只是 的一个算例。
- 算例参数下的理论上限是 ~211 TFLOPs(仍带宽 bound):这是按 、K/V 重读全落 HBM 估的保守值;FA3 用更大的块和 Hopper 特性,实测吞吐在这个量级之上(第 17 章)。
- 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 论文口径),手写骨架只会更低。
本章动手练习:
- 用 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 读写流量,验证差距。- 推导 FA 反向的 HBM 流量公式,对比朴素反向(需要 N×N 中间梯度矩阵)。
- 把 FA1 论文 Algorithm 1(外 K/V 内 Q)和 FA2 论文 Algorithm 1(外 Q 内 K/V)并排放,逐行找出两者的差异——尤其是 、、O 分别在哪一层循环里活着。