一句话总结:FlashAttention 不是近似 Attention,也没有把 N2N^2 次配对计算变没。它通过分块和 Online Softmax,避免把完整的 N×NN \times N 中间矩阵反复写入显存。

朴素 Attention 把完整分数和概率矩阵写回 HBM,FlashAttention 在片上分块合并结果

21.1 瓶颈不只是“算得多”

对一个 Attention head,设 Q,K,V∈RN×dQ,K,V \in \mathbb{R}^{N \times d}:

O=softmax⁡ ⁣(QKTd+Bmask)VO = \operatorname{softmax}\!\left(\frac{QK^T}{\sqrt d} + B_{\mathrm{mask}}\right)V

计算量的主项仍是 Θ(N2d)\Theta(N^2d)。朴素实现还会 materialize 一个 N×NN \times N 的分数矩阵,Softmax 后可能再 materialize 一个同样大的概率矩阵。

当 N=4096N=4096 时,单个 head、单个样本的一张 FP16 分数矩阵有:

40962×2 bytes=32 MiB4096^2 \times 2\ \text{bytes} = 32\ \text{MiB}

如果朴素路径在 batch size 8、32 heads 上同时保留这一张矩阵,就是 8 GiB。这还没计 mask、Softmax 输出、Dropout 与反向传播所需的其他状态。这是对朴素 materialized 路径的算账,不是说所有现代 PyTorch Attention 都会真的分配这 8 GiB。

问题的另一半是数据搬运:

HBM 读 Q/K → 算 S=QKᵀ → 把 S 写回 HBM
HBM 读 S   → Mask/Softmax → 把 P 写回 HBM
HBM 读 P/V → 算 O=PV → 把 O 写回 HBM

GPU 的 Tensor Cores 可以很快地做矩阵乘,但计算单元不工作时也要等数据。FlashAttention 从这个 I/O 瓶颈入手。


21.2 分块:大纸不落地,小块在桌上算完

还用原来的书桌类比:HBM 像容量很大的书架,片上 registers/shared memory/L1 像面积很小的桌面。FlashAttention 每次只把 Qi,Kj,VjQ_i,K_j,V_j 的 tile 搬上桌:

Sᵢⱼ = QᵢKⱼᵀ / √d + mask
                 ↓
        Online Softmax 统计量
                 ↓
          累加到输出 Oᵢ
Q tile 与 K/V tile 在片上计算局部分数,完整 N 乘 N 矩阵不落到 HBM

这里容易把几个数字混在一起。以 A100 为例:整张 GPU 有 40 MB L2;每个 SM 有最大 192 KB 的 L1/texture/shared-memory 组合数据 cache,其中 shared-memory carveout 还要按 kernel 配置。它们不是一个可以随便合并使用的“20 MB SRAM 池”。A100 40GB 的 HBM2 带宽约 1550 GB/s,80GB 版约 2039 GB/s;硬件版本不同,不该用一个“SRAM 永远快 20 倍”的表概括。

原始论文 Algorithm 1 的理想化 tile 大小是:

Bc=⌈M4d⌉,Br=min⁡ ⁣(⌈M4d⌉,d)B_c=\left\lceil\frac{M}{4d}\right\rceil,\qquad B_r=\min\!\left(\left\lceil\frac{M}{4d}\right\rceil,d\right)

其中 dd 是每个 head 的维度,MM 是论文抽象机里片上可容纳的标量元素数,不是把 192 KB 字节数不管 dtype 直接代入。真实 kernel 还要考虑 dtype、register pressure、shared-memory carveout、alignment、warp 分工与硬件世代,不会从这一条式子唯一推出“tile 必须是 64”。


21.3 Online Softmax:看一块,仍然得到整行的结果

Softmax 的分母需要整行,所以不能对每个 tile 各做一次 Softmax 再直接拼起来。对某一行 query,已处理过的 key blocks 维护:

  • mm:当前最大 score;
  • ℓ\ell:在当前最大值尺度下的指数和;
  • aa:尚未除以 ℓ\ell 的 Value 加权和。

新 score block ss 和对应 VjV_j 到来时:

m′=max⁡(m,max⁡(s)),α=em−m′m'=\max(m,\max(s)),\qquad \alpha=e^{m-m'}
p=es−m′,ℓ′=αℓ+∑p,a′=αa+pVjp=e^{s-m'},\qquad \ell'=\alpha\ell+\sum p,\qquad a'=\alpha a+pV_j

最后 o=a/ℓo=a/\ell。如果后面的 tile 出现更大的 score,α\alpha 会把之前的累积量换到新的数值尺度上。

Online Softmax 用运行最大值、分母和未归一化输出合并多个 score blocks

下面是不带 mask 的 PyTorch 教学版。它不是高性能 kernel,但可以验证合并公式:

import math
import torch


def online_attention(q, k, v, block_size):
    # q: [Nq, d], k: [Nk, d], v: [Nk, dv]
    scale = 1.0 / math.sqrt(q.size(-1))
    rows = q.size(0)
    m = torch.full((rows,), -torch.inf, device=q.device, dtype=q.dtype)
    denominator = torch.zeros_like(m)
    accumulator = torch.zeros(
        rows, v.size(-1), device=q.device, dtype=q.dtype
    )

    for start in range(0, k.size(0), block_size):
        k_block = k[start:start + block_size]
        v_block = v[start:start + block_size]
        scores = q @ k_block.transpose(-2, -1) * scale

        block_max = scores.max(dim=-1).values
        new_m = torch.maximum(m, block_max)
        correction = torch.exp(m - new_m)
        probabilities = torch.exp(scores - new_m[:, None])

        denominator = (
            correction * denominator + probabilities.sum(dim=-1)
        )
        accumulator = (
            correction[:, None] * accumulator + probabilities @ v_block
        )
        m = new_m

    return accumulator / denominator[:, None]

对同一组 Q,K,VQ,K,V,它应该与直接计算 softmax(Q @ K.T / sqrt(d)) @ V 在浮点误差内相等。“Exact Attention”指不改变数学 Attention、不做稀疏或低秩近似;计算顺序改变后,不承诺浮点结果 bit-for-bit 相同。


21.4 别把计算量、额外内存和 HBM I/O 混在一起

FlashAttention 不改变 N 平方的 Attention 配对计算,但把额外内存从二次降为线性并降低 HBM 访问
问题朴素 materialized AttentionFlashAttention
算术主项Θ(N2d)\Theta(N^2d)Θ(N2d)\Theta(N^2d)
额外中间内存(固定 d)Θ(N2)\Theta(N^2)O(N)O(N)
原始论文抽象机上的 HBM 访问Θ(Nd+N2)\Theta(Nd+N^2)Θ(N2d2/M)\Theta(N^2d^2/M)

第三行有明确前提:MM 是片上容量(按元素计),且论文分析的范围是 d≤M≤Ndd \le M \le Nd。它不能随手化简成 O(N2d/M)O(N^2d/M),也不能从“二次变线性”就断言实际显存一定省 N 倍。真实显存还有 Q/K/V、输出、模型权重、optimizer state、allocator 和其他层。

反向传播时,FlashAttention 不保存完整概率矩阵,而是利用保留的统计量重算局部 tile。这确实增加了一些 FLOPs,但不能因此笼统说“反向一定更慢”;少了 HBM 往返后,kernel 和端到端训练都可能更快。


21.5 FA1 到 FA4:核心算法延续,kernel 跟着硬件变

截至 2026 年 7 月 FlashAttention 1 到 4 的算法和硬件重点
  • FA1(2022):确立 I/O-aware tiling + Online Softmax,避免 materialize 完整矩阵。
  • FA2(2023):减少非 matmul FLOPs,在 sequence 维度增加 thread-block 并行,改善 warp 分工。FA2 论文在 A100 的特定 benchmark 中报告约 2× FA1、50–73% 理论峰值,GPT-style 训练最高 225 TFLOPs/s/GPU(72% model FLOPs utilization)。这些是论文实验边界,不是任意模型的保底加速。
  • FA3(2024,官方仓库仍标为 Hopper beta):利用 H100/H800 的 TMA、WGMMA 与 FP8 路径。
  • FA4(2026):针对 Blackwell B200/GB200 的不对称硬件扩展,用 CuTe DSL 重做 pipeline 协同设计。截至 2026-07,官方仓库提供 flash-attn-4 pre-release;具体 head dimension、dtype 与 backward 能力仍要查当前 release,不应从 FA2 规则类推。

版本号不是“越新就在所有 GPU 上越快”。FA3 针对 Hopper,FA4 针对 Blackwell;Ampere/Ada 仍可能使用 FA2。


21.6 实际代码:先让框架调度,再强制 kernel

PyTorch 的 scaled_dot_product_attention 会根据 device、dtype、shape、mask 等条件在可用 backend 中选择。对本书第 18 章的 tensor layout,Q/K/V 先整理为 [batch, heads, time, head_dim]:

import torch.nn.functional as F


def attention(q, k, v, dropout_p, training):
    return F.scaled_dot_product_attention(
        q,
        k,
        v,
        attn_mask=None,
        dropout_p=dropout_p if training else 0.0,
        is_causal=True,
    )

scaled_dot_product_attention 会按传入的 dropout_p 始终做 Dropout,它不会自动读取外层 module 的 training 状态,所以 eval 时必须显式传 0.0。

需要证明某个输入确实走 Flash backend 时,可以在诊断代码里限定 backend:

from torch.nn.attention import SDPBackend, sdpa_kernel

with sdpa_kernel(SDPBackend.FLASH_ATTENTION):
    output = F.scaled_dot_product_attention(
        q, k, v, dropout_p=0.0, is_causal=True
    )

如果 fused kernel 不支持当前输入,PyTorch 会给出原因;生产代码通常不应为了“名字上用了 Flash”而禁掉所有 fallback。

官方 Dao-AILab/flash-attention 包也可以直接调用。但硬件范围已经不是旧版的“只支持新 NVIDIA”:截至 2026-07,FA2 CUDA 路径覆盖 Ampere/Ada/Hopper,官方仓库也提供 ROCm 的 CK 与 Triton AMD backends。安装前查当前 README 的 PyTorch、CUDA/ROCm、GPU、dtype 和 head-dimension 矩阵,比在书里永久写死一条 pip install 更可靠。


21.7 怎样诚实地做 benchmark?

至少固定并报告:

  • GPU 型号与功耗模式,PyTorch/CUDA/ROCm/kernel 版本;
  • dtype,batch、heads、query length、key length、head dimension;
  • causal/非 causal,mask、Dropout,forward 还是 forward+backward;
  • warmup 次数、计时方法、median/percentile,是 kernel microbenchmark 还是整模型 wall time;
  • 峰值显存与输出/梯度误差。

短序列或小 batch 时,kernel launch 和 layout 转换可能盖过节省;某个 Attention kernel 快 2×,也不等于整个模型或每 token latency 快 2×。

FlashAttention 与 KV Cache 也解决不同问题:前者降低一次 Attention的 I/O 代价;后者避免 autoregressive decoding 步之间重算旧 K/V。


21.8 本章总结

  • FlashAttention 保留标准 Attention 的数学结果,核心是 I/O-aware 实现;
  • 它仍做 Θ(N2d)\Theta(N^2d) 的主要算术,却不 materialize 完整 N×NN \times N 中间矩阵;
  • Online Softmax 同时重标定分母与 Value 加权和,所以 tile 结果能精确合并;
  • dd 是 head dimension,MM 是抽象片上容量;不能用模型宽度和字节数乱代入 tile 式子;
  • FA1–FA4 的 kernel 与硬件范围不断变化,速度数字必须带 workload 和版本边界;
  • PyTorch SDPA 能自动调度 backend,但 eval 时 dropout_p=0.0 仍是调用者的责任。

本章交付物

  • 能分开计算复杂度、额外内存复杂度与 HBM I/O
  • 能用 m,ℓ,am,\ell,a 推导 Online Softmax 的 tile 合并
  • 能解释为什么“exact”不等于 bit-for-bit 相同
  • 能在 PyTorch 中正确调用 SDPA,并验证实际 backend
  • 能设计一个不把 microbenchmark 冒充整模型速度的测试

一手资料


下一章预告

FlashAttention 让当前这一次 Attention 更高效,但第 20 章的自回归循环还在每一步重算旧 token 的 K 和 V。第 22 章用 KV Cache 解决这个跨 step 冗余。

引用本文 / Cite
Zhang, Wayland (2026). 第 21 章:Flash Attention - 内存优化原理. In Transformer 架构:从直觉到实现. https://waylandz.com/llm-transformer-book/%E7%AC%AC21%E7%AB%A0-Flash-Attention-%E5%86%85%E5%AD%98%E4%BC%98%E5%8C%96%E5%8E%9F%E7%90%86/
@incollection{zhang2026transformer_21_-Flash-Attention-,
  author = {Zhang, Wayland},
  title = {第 21 章:Flash Attention - 内存优化原理},
  booktitle = {Transformer 架构:从直觉到实现},
  year = {2026},
  url = {https://waylandz.com/llm-transformer-book/%E7%AC%AC21%E7%AB%A0-Flash-Attention-%E5%86%85%E5%AD%98%E4%BC%98%E5%8C%96%E5%8E%9F%E7%90%86/}
}