Skip to content
kefan.life
Go back

FlashAttention 详解(V1 & V2)

标准 Attention 的计算流程是 S=QKS = QK^\topP=softmax(S)P = \text{softmax}(S)O=PVO = PV,其中 Q,K,VRN×dQ, K, V \in \mathbb{R}^{N \times d}NN 为序列长度,dd 为头维度)。两次矩阵乘法的总 FLOPs 为 4N2d4N^2d,但硬件真实瓶颈往往不在计算量,而在中间矩阵 SSPP:它们各占 N×NN \times N 的显存,必须完整写入 HBM 再读出。这不仅导致 IO 成为性能瓶颈(算术强度1d/264d/2 \approx 64 FLOPs/Byte,远低于 A100 拐点 156),更直接带来显存容量问题:N=4096N = 4096、32 个头时,仅 SSPP 就要占用约 2 GB 显存,序列再长就直接 OOM。

FlashAttention 的优化方向因此是减少 HBM 访存,同时消除 N×NN \times N 矩阵的显存占用。做法是两个技术的配合:

本文从这两个基础技术出发,逐步拆解 FlashAttention V1 和 V2 的计算过程。如果你对 Attention 基本原理还不熟悉,建议先阅读《为什么 KV 缓存没有 Q》

FlashAttention 系列总览

FlashAttention 从 2022 年至今迭代了三个版本,核心思路一脉相承:将 Attention 计算拆成小块,在 SRAM 中完成,避免 N×NN \times N 中间矩阵落地 HBM。

V1245(Dao et al., 2022)引入 Tiling + Online Softmax

FlashAttention 核心概念:左侧为 GPU 存储层级(SRAM 带宽 19 TB/s 但容量仅 20 MB,HBM 带宽 1.5 TB/s 但容量 40 GB),中间为 Tiling 计算流程,右侧为 GPT-2 上的性能对比。标准实现中 Softmax、Dropout、Mask 等非矩阵运算占据了大部分时间,FlashAttention 的融合内核将总耗时缩短至约 1/7(图源)

V267(Dao, 2023)做了三项工程优化,将 A100 利用率从 25-40% 提升到 50-73%:

V38(Shah, Dao et al., 2024)针对 H100 Hopper 架构做硬件级适配:

下面先给出 Online Softmax 和 Tiling 的公式框架,再进入 V1 和 V2 的完整算法。

两个核心技术

Online Softmax:流式精确归一化

标准 Safe Softmax 对输入向量 xx(在 Attention 语境下即 SS 的某一行)需要三遍扫描:找全局最大值 mm(避免后续指数溢出),算指数和 =exim\ell = \sum e^{x_i - m},算归一化输出。三遍之间有严格的数据依赖。分块计算时无法提前获得全局 mm,Online Softmax3 的解法是边扫描边用递推公式动态修正。

核心恒等式:

exmnew=exmoldemoldmnewe^{x - m_{\text{new}}} = e^{x - m_{\text{old}}} \cdot e^{m_{\text{old}} - m_{\text{new}}}

发现更大的 mnewm_{\text{new}} 时,之前以 moldm_{\text{old}} 为基准的指数项乘以修正因子 emoldmnewe^{m_{\text{old}} - m_{\text{new}}} 即可对齐。纯代数恒等变形,没有近似。

KK 被分成 TcT_c 块,上标 (j)(j) 表示处理完前 jj 个块后的累积状态。当第 jj 块到来时:

m(j)=max(m(j1),  max(x(j)))m^{(j)} = \max\left(m^{(j-1)},\; \max(x^{(j)})\right) (j)=(j1)em(j1)m(j)+iexi(j)m(j)\ell^{(j)} = \ell^{(j-1)} \cdot e^{m^{(j-1)} - m^{(j)}} + \sum_i e^{x_i^{(j)} - m^{(j)}}

从 Softmax 到 Attention Output

Attention 最终需要的不是 Softmax 概率本身,而是 O=softmax(S)VO = \text{softmax}(S) \cdot V。Online Softmax 的递推直接扩展到输出 OO 上。设 S(j)S^{(j)}V(j)V^{(j)} 分别为第 jj 个 K/V 块对应的注意力分数和 Value,处理完前 jj 个块后,未归一化的输出累加器为:

O~(j)=em(j1)m(j)O~(j1)+eS(j)m(j)V(j)\tilde{O}^{(j)} = e^{m^{(j-1)} - m^{(j)}} \cdot \tilde{O}^{(j-1)} + e^{S^{(j)} - m^{(j)}} \cdot V^{(j)}

所有块处理完毕后一次性归一化:O=O~(Tc)/(Tc)O = \tilde{O}^{(T_c)} / \ell^{(T_c)}

FlashAttention 前向传播中的分块与 rescaling 机制(图中省略了 max 减法以简化展示):Q 与 K 的不同块分别计算局部 S 和 A,在 SRAM 中完成指数运算,不落地 HBM。输出 O^{(2)} 通过 rescaling 因子 \ell^{(1)} / \ell^{(2)} 修正 O^{(1)},再加上新块的贡献(来自 Tri Dao 的博客)

Tiling:分块驻留 SRAM

QQ 沿序列维度分成 Tr=N/BrT_r = \lceil N / B_r \rceil 块,K,VK, V 分成 Tc=N/BcT_c = \lceil N / B_c \rceil 块。块大小由 SRAM 容量 MM(元素数)和头维度 dd 决定:

Bc=M4d,Br=min(M4d,  d)B_c = \left\lceil \frac{M}{4d} \right\rceil, \quad B_r = \min\left(\left\lceil \frac{M}{4d} \right\rceil,\; d\right)

分母 4 可以理解为论文在 SRAM 预算下给出的保守容量约束:片上主要部分是 Q/K/V/OQ/K/V/O 四类 tile。实际实现里 BrB_rBcB_c 可以不同,也会受中间状态、统计量和硬件资源占用影响。

FlashAttention V1 的分块策略:Q、K、V 沿序列维度切分为行块,外循环固定 K/V 块 j,内循环遍历 Q 块 i。每次迭代将 Q_i, K_j, V_j 加载到 SRAM 中,在片上计算局部注意力分数 S_{ij} = Q_i K_j^\top,通过 Online Softmax 更新累加输出

FlashAttention V1:算法详解

前向传播

V1 的前向传播是一个双层循环,外循环遍历 K/V 块(jj),内循环遍历 Q 块(ii):

初始化 O = 0, m = -∞, l = 0

外循环:for j = 1 to T_c
    从 HBM 加载 K_j, V_j 到 SRAM(常驻至内循环结束)

    内循环:for i = 1 to T_r
        从 HBM 加载 Q_i, O_i, m_i, l_i 到 SRAM
        计算 S_ij = Q_i · K_j^T / √d          ← SRAM 内完成
        更新 m, l, O(Online Softmax)          ← SRAM 内完成
        将 O_i, m_i, l_i 写回 HBM

每次内循环迭代的计算步骤:

Step 1 — 局部注意力分数:Sij=QiKj/dRBr×BcS_{ij} = Q_i K_j^\top / \sqrt{d} \in \mathbb{R}^{B_r \times B_c}

Step 2 — 更新行最大值:minew=max(mi,  rowmax(Sij))m_i^{\text{new}} = \max\left(m_i, \;\text{rowmax}(S_{ij})\right)

Step 3 — 更新指数和:inew=iemiminew+rowsum(eSijminew)\ell_i^{\text{new}} = \ell_i \cdot e^{m_i - m_i^{\text{new}}} + \text{rowsum}\left(e^{S_{ij} - m_i^{\text{new}}}\right)

Step 4 — 更新输出:

Oinew=Oiiemiminewinew+eSijminewinewVjO_i^{\text{new}} = O_i \cdot \frac{\ell_i \cdot e^{m_i - m_i^{\text{new}}}}{\ell_i^{\text{new}}} + \frac{e^{S_{ij} - m_i^{\text{new}}}}{\ell_i^{\text{new}}} \cdot V_j

几个细节:

mm 初始化为 -\infty\ell 初始化为 00 第一块数据进来时 e=0e^{-\infty} = 0,修正项自动清零,公式退化为普通的局部计算。用数学边界值替代 if (is_first_iteration) 分支判断。

mm\ell 是向量(RBr\mathbb{R}^{B_r})而非标量。 Softmax 按行独立归一化:SS 的每一行对应一个 Query token,各有自己的注意力分布。BrB_r 个 token 就需要 BrB_r 个独立的 mm\ell

数值稳定性。 所有指数输入都减去了当前最大值,保证 x0x \le 0。修正因子 emoldmnew1e^{m_{\text{old}} - m_{\text{new}}} \le 1,不会溢出。

CUDA 执行模型

V1 以 grid = (Batch, Head) 启动内核:每个 (batch, head) 对应一个 thread block,被调度到某个 SM 上执行(一个 SM 可同时驻留多个 thread block,具体取决于资源占用),串行跑完双层循环。

瓶颈:并行度绑定在 batch size ×\times head 数上。batch=1、head=32 时,A100 的 108 个 SM 只有 32 个在工作。这是 V2 翻转循环顺序的动因。

反向传播与重计算

FlashAttention 不存储 SSPP,前向只保存 Q,K,V,O,m,Q, K, V, O, m, \ell,反向时在 SRAM 中重计算:

Pij=diag(i)1eSijmiP_{ij} = \text{diag}(\ell_i)^{-1} \cdot e^{S_{ij} - m_i}

重计算增加约 50% 前向 FLOPs,但由于标准 Attention 是 memory-bound 的,额外计算时间远小于省下的 IO 时间。以 A100 处理 N=4096,d=128N = 4096, d = 128 的单头为例:

项目A100 耗时
重计算 QKQK^\top2×40962×1284.3×1092 \times 4096^2 \times 128 \approx 4.3 \times 10^9 FLOPs14  μs\approx 14\;\mu s
从 HBM 读 SS 矩阵40962×233.54096^2 \times 2 \approx 33.5 MB17  μs\approx 17\;\mu s

片上重算比从 HBM 读一遍还快,而标准反向传播涉及 S,PS, P 的多次读写。重计算用微量计算时间换下大量 IO 等待,墙钟时间反而缩短。

IO 复杂度

标准 Attention 的 HBM 访问量为 O(Nd+N2)O(Nd + N^2)N2N^2 项来自中间矩阵 S,PS, P

FlashAttention V1 的 HBM 访问量为 O(N2d2/M)O(N^2d^2/M)2。推导:外循环 Tc=O(Nd/M)T_c = O(Nd/M) 次迭代,每次读取全部 Q,OQ, O(共 O(Nd)O(Nd)),总计 TcNd=O(N2d2/M)T_c \cdot Nd = O(N^2d^2/M)。由于 d2Md^2 \ll Md=128d = 128d2=16384d^2 = 16384,A100 SRAM 约 9830498304 个 FP16 元素),这远小于 O(N2)O(N^2)

精确 Attention 的 HBM 访问下界为 Ω(N2d2/M)\Omega(N^2d^2/M),因此 FlashAttention 的 IO 复杂度是渐近最优的2

FlashAttention V2:三项改进

V1 在 A100 上只达到 25-40% 理论算力利用率。V2 将其提升到 50-73%,比 V1 快约 1.7x6

翻转循环顺序

V1 外循环 K/V、内循环 Q,每次外循环切换 KjK_j 时所有 Q 块的统计量都要重新从 HBM 加载再写回,OO 被读写 2Tc2T_c 次。

V2 反转:外循环 Q 块(映射到不同 thread block 并行),内循环 K/V 块

外循环(并行):for i = 1 to T_r
    Q_i 加载到 SRAM(全程常驻)
    初始化 m_i = -∞, l_i = 0, Õ_i = 0

    内循环:for j = 1 to T_c
        从 HBM 加载 K_j, V_j 到 SRAM
        计算 S_ij,更新 m, l, Õ(全在 SRAM 中)

    O_i = Õ_i / l_i              ← 最终归一化
    L_i = m_i + log(l_i)         ← 保存 logsumexp
    将 O_i, L_i 写回 HBM         ← 只写一次

三个收益:OO 只写 HBM 一次;Q 全程驻留 SRAM;不同 Q 块之间无递推依赖,Grid 从 (Batch, Head) 扩展为 (Batch, Head, T_r)

延迟归一化

V1 每步维护归一化后的 OiO_i,需要先反归一化(乘回旧分母)再重归一化(除以新分母),每步 3 次标量运算。V2 改为维护未归一化的 O~\tilde{O}

O~i(j)=emi(j1)mi(j)O~i(j1)+eSijmi(j)Vj\tilde{O}_i^{(j)} = e^{m_i^{(j-1)} - m_i^{(j)}} \cdot \tilde{O}_i^{(j-1)} + e^{S_{ij} - m_i^{(j)}} \cdot V_j

最后一步才归一化:Oi=O~i(Tc)/i(Tc)O_i = \tilde{O}_i^{(T_c)} / \ell_i^{(T_c)}。这些逐元素运算跑在 CUDA Core 上,而矩阵乘法跑在 Tensor Core 上。A100 的 FP16/BF16 matmul 峰值约 312 TFLOPS,non-matmul FP32 只有 19.5 TFLOPS,单个 non-matmul FLOP 约贵 16 倍6。减少非 GEMM 操作直接降低了这部分的时间占比。

Warp 分工沿 Q 维度切分

V1 沿 K/V 维度切分 warp,每个 warp 只拿到 SijS_{ij} 的部分列,softmax 的 rowmax/rowsum 必须跨 warp 同步合并。V2 沿 Q 维度切分,每个 warp 独占自己的行,无需通信:

V1(沿 K/V 切分)V2(沿 Q 切分)
Softmax 统计部分列,跨 warp 合并完整行,独立完成
输出 OO 写入多 warp 更新同一行各 warp 独占不同行
通信开销shared memory / shuffle

Causal Mask 优化

因果注意力要求位置 ii 只关注 jij \le i。V2 在块级别判断:完全在对角线以下的块正常算,完全在对角线以上的块直接跳过,跨对角线的块逐元素 mask。对角线以上约占一半,块级跳过使因果注意力实际开销接近非因果的一半。

LogSumExp 合并统计量

V1 向 HBM 写入 mim_ii\ell_i 两个向量。V2 将它们合并为一个标量 Li=mi+log(i)L_i = m_i + \log(\ell_i),存储减半。

为什么一个值就够?回忆 Softmax 的定义:

Pij=eSijmiiP_{ij} = \frac{e^{S_{ij} - m_i}}{\ell_i}

取对数:logPij=Sijmilogi=Sij(mi+logi)=SijLi\log P_{ij} = S_{ij} - m_i - \log \ell_i = S_{ij} - (m_i + \log \ell_i) = S_{ij} - L_i

因此 Pij=eSijLiP_{ij} = e^{S_{ij} - L_i}。反向传播需要恢复 PijP_{ij} 时,只需从 HBM 读取 LiL_i(而非分别读 mim_ii\ell_i),配合重计算得到的 SijS_{ij},一步指数即可还原。LiL_i 本质上就是 log-normalizer:Li=logjeSijL_i = \log \sum_j e^{S_{ij}}(数值稳定形式为 mi+logim_i + \log \ell_i)。

性能对比6

序列长度V1 TFLOPSV2 TFLOPSV2 利用率
102412419663%
204813621870%
409614122773%
819213822271%

结语

FlashAttention 的核心原则是:在 GPU 算力远超带宽的硬件现实下,用计算换 IO。 Tiling 把数据锁在 SRAM 里,Online Softmax 在不知道全局统计量时实现精确归一化,重计算替代存储。每一步都在将算子从 memory-bound 推向 compute-bound。

从 V1 到 V3 的迭代展现了系统优化的层级递进:V1 解决算法问题(如何正确分块),V2 解决工程问题(如何利用硬件并行度),V3 解决硬件适配问题(如何压榨特定架构的每一条执行通道)。层次在变,原则不变:找到硬件的真实瓶颈,然后针对性地设计算法。

参考资料

  1. Roofline: An Insightful Visual Performance Model for Multicore Architectures (Williams et al., 2009)

  2. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness (Dao et al., 2022) 2 3 4 5 6 7

  3. Online Normalizer Calculation for Softmax (Milakov & Gimelshein, 2018) 2

  4. FlashAttention V1 详解 — AIInfraGuide

  5. ELI5: FlashAttention (Aleksa Gordić)

  6. FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning (Dao, 2023) 2 3 4

  7. FlashAttention V2 详解 — AIInfraGuide

  8. FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision (Shah, Dao et al., 2024) 2


Share this post on:

Previous Post
重新理解 MLA:缓存 X 而非 KV
Next Post
无损投机采样的概率论视角