Skip to content

FlashAttention 详解:从 Online Softmax 到单遍融合

注意力机制是 Transformer 模型的核心组件。所有主流模型架构,如 GPT、LLaMA 和混合专家(Mixture of Experts, MoE),都依赖它来连接 Token 并构建语义。

但注意力是昂贵的。它的计算涉及大规模矩阵乘法,更昂贵的是 GPU 内存与计算单元之间海量的数据移动。随着序列长度不断增长,注意力的标准实现面临显存与计算效率的双重瓶颈。

1. 为什么朴素实现很慢?

回顾注意力运算的定义:

X=QK,A=softmax(X),O=AV

对长度为 N、头维度为 d 的序列,查询矩阵 Q、键矩阵 K、值矩阵 V 的大小都是 N×d,而注意力分数矩阵 X 与概率矩阵 A 的大小都是 N×N

朴素实现慢的根源在于 GPU 的内存层级。计算核心只能直接使用一小块非常快的片上内存(SRAM),而更大但更慢的内存池(HBM,即高带宽内存)位于片外。深度学习模型中的大多数操作都是“内存受限”的:它们的速度受限于在 HBM 与 SRAM 之间搬运数据所需的时间,而非在 SRAM 中执行的算术计算。推荐阅读 Horace He 的博客,快速了解深度学习中计算、内存与开销的概念。

FlashAttention 论文中的示意图

注意力的标准实现会在 HBM 中物化 N×N 的矩阵 XA。这涉及:

  1. 从 HBM 读取 QK
  2. 计算 X=QK,把 X 写回 HBM。
  3. 从 HBM 读取 X
  4. 计算 A=softmax(X),把 A 写回 HBM。
  5. 从 HBM 读取 AV
  6. 计算 O=AV,把 O 写回 HBM。

对于长序列(Nd),XAO(N2) 次内存读写主导了运行时间。

问题出在 HBM 流量上。设每个元素占 w 字节:实现先写后读 N×N 矩阵 X(共 2N2w 字节),然后再写读 N×N 矩阵 A(又 2N2w 字节)。这四次在庞大而缓慢的 HBM 之间的来回搬运,压过了读取初始输入 Q,K,V 和写入最终输出 OO(Ndw) 流量。因此总 HBM 访问量为 O(N2+Nd) 个元素,在长序列下即 O(N2) 量级。

2. FlashAttention 是什么?

FlashAttention 是一种经过优化的 Transformer 注意力机制,它让注意力在 GPU 上的运行速度大幅提升,同时更加节省内存。

标准注意力与 FlashAttention 对比

GPU 有两种主要的内存类型。高带宽内存(HBM)容量大但速度相对较慢;片上 SRAM 速度极快但容量非常有限。

标准自注意力在这两者之间不断地搬运数据。这种来回搬运代价高昂,并且随着序列长度的增长,其成本变得十分显著。

FlashAttention 通过将注意力计算拆分为完全能放进快速 SRAM 的小块(Tile)来规避这一问题。每个小块被端到端地完整处理,Softmax 以增量方式应用,因此中间结果无需写回 HBM。由此,完整的注意力矩阵永远不会存储在内存中。

与稀疏注意力或线性注意力方法不同,FlashAttention 不是一种近似方法。它产生的数学输出与标准自注意力完全相同,只是以一种更节省内存的方式执行。

3. FlashAttention 是如何工作的?

FlashAttention 通过重新设计注意力在 GPU 上的计算方式来实现其高效性。它遵循一个简单的机制:在快速片上内存中尽可能多地完成计算,并避免不必要的慢速内存移动。

一个有助于理解的比喻是厨房。GPU 的片上 SRAM 就像一张小而快的厨房台面,是真正准备和烹饪的地方;GPU 高带宽内存(HBM)就像街角的大型杂货店,可以存储你需要的所有东西,但来回往返需要时间。

简单来说,标准注意力每完成一步就要跑到杂货店一趟。相比之下,FlashAttention 会规划好烹饪流程,让你烹饪时所有东西都摆在台面上。

FlashAttention 工作机制

FlashAttention 依赖两个关键思想:分块(Tiling)与重计算(Recomputation)。

3.1 分块(Tiling)

继续用烹饪的例子:分块就是 FlashAttention 把注意力计算塞进小台面的方法。

FlashAttention 不是加载整个序列、构建完整的注意力矩阵,而是把输入拆成小块(Tile)。每个块都能完全放进 GPU 的快速 SRAM 中。FlashAttention 一次处理一个块,从开始到结束,然后才进入下一个块。

回到厨房类比:你无法把一整场宴会的食材都放在小台面上,所以你要小批量地备菜和烹饪——切几样蔬菜、炒熟、清空台面,再处理下一批。通过这种方式,你避免了不断往返杂货店。

这种逐块执行的方式让 FlashAttention 把数据保持在本地,既快又高效,而且永远不会在慢速内存中物化出完整的注意力矩阵。

FlashAttention 中的分块

计算被拆成若干块。矩阵 Q,K,V 被划分为更小的子矩阵。算法遍历 KV 的各个分块,把它们加载进 SRAM。在这个外层循环中,它再遍历 Q 的各个分块。

循环的简化视图:

text
// O 是最终输出矩阵,初始化为零
// l, m 是在线 Softmax 的统计量,已初始化
for each block K_j, V_j in K, V:
  Load K_j, V_j into fast SRAM
  for each block Q_i in Q:
    Load Q_i, O_i, l_i, m_i into fast SRAM
    // 核心计算,全部在片上完成
    Compute S_ij = Q_i @ K_j^T
    Compute P_ij = softmax(S_ij)  //(这是简化写法,在线 Softmax 的推导见第 4 节)
    Compute O_ij = P_ij @ V_j
    // 用新结果更新输出分块 O_i
    Update O_i using O_ij and softmax statistics
    // 把更新后的 O_i, l_i, m_i 写回 HBM
    Write O_i, l_i, m_i to HBM

这种分块结构改变了内存访问模式。算法分块遍历 KV,每个元素只从 HBM 读取一次。关键在于内层循环:对每个加载进 SRAM 的 Kj 分块,算法必须遍历 Q 的所有分块,才能计算出对输出的相应更新。这意味着算法会对 Q 矩阵做多遍遍历。遍历的次数由 K 的分块数决定,即 T=N/b,其中 b 是分块大小。由于大小为 M 的片上 SRAM 必须容纳一个 K 分块(大小 b×d),分块大小 b 被限制在 O(M/d)。这导致 T=O(Nd/M) 遍遍历。每一遍都会读取完整的 Q(并读写 O 的分块),从而产生 O(Nd×T)=O(N2d2/M) 的 HBM 流量。这个数量远小于标准注意力的 O(N2) 流量。

示例:为了具体化,考虑一块 NVIDIA A100 GPU。它的每个流式多处理器(SM)有 192 KB SRAM。单个 Kernel 会使用其中的一部分,假设工作 SRAM 大小为 M_bytes = 128 KB。如果使用 bfloat16 精度(每个数字 2 字节),有效 SRAM 大小按元素计为 M = 128 * 1024 / 2 = 65,536 个元素。对典型的头维度 d = 64d2=4096,内存访问的缩减因子约为 M/d2=65536/4096=16,即 FlashAttention 的 HBM 访问次数约为标准实现的 1/16。对 d = 128,因子为 M/d2=65536/16384=4

3.2 重计算(Recomputation)

在训练过程中,标准注意力会存储大量中间结果,以便在反向传播时复用。这种存储带来了高昂的内存代价。FlashAttention 采取了不同的做法:它不存储这些中间结果,而是在需要时重新计算注意力分数的小部分。

回到厨房,这就像切洋葱。你可以走去杂货店把切好的洋葱存起来,之后再走回去取;也可以把切好的扔了,等烹饪时重新切一颗新鲜的。令人惊讶的是,第二种方式反而更快,因为它避免了频繁或较长的往返。

在现代 GPU 上,重计算遵循同样的逻辑,因为与内存移动相比,额外的计算非常廉价。通过重新计算小值而不是存储和加载它们,FlashAttention 显著减少了内存流量,同时保持了训练的高效性。

分块与重计算共同作用,让 FlashAttention 把注意力计算留在“台面”上,最小化“去杂货店”的次数,并充分利用现代 GPU 硬件的优势。

4. FlashAttention 数学推导:从 Online Softmax 到单遍融合

FlashAttention [1] 的关键创新,是借鉴 Online Softmax [3] 的思想对自注意力计算进行分块(Tiling),从而把整个多头注意力层融合到单个 Kernel 中,无需把中间的 Logits 和注意力分数写入 GPU 全局内存。接下来,我们从 Online Softmax 讲起,逐步推导 FlashAttention 如何通过分块(Tiling)与在线融合,将注意力计算压缩为单轮遍历,并尽量在 GPU 片上内存(SRAM)中完成,从而避免 O(N2) 的中间显存开销。

4.1 标准自注意力

忽略注意力头(head)与批次(batch)维度(这两个维度上的计算完全并行),同时省略注意力掩码(mask)和缩放因子 1d 等细节,标准自注意力的计算可以概括为:

(1)O=softmax(QK)V

其中 Q,K,V,ORN×dN 为序列长度,d 为每个注意力头的维度,Softmax 作用于最后一个维度(列方向)。

标准的计算方式是把自注意力分解为三步:

X=QK,A=softmax(X),O=AV

其中 X 为 Pre-softmax Logits,A 为注意力分数矩阵,O 为输出矩阵。

直接实现需要把完整的 N×N 中间矩阵 XA 写入显存再读回,显存开销与访问量均随序列长度平方增长。

FlashAttention 的关键在于:它无需在全局内存中物化 XA 矩阵,而是把公式 (1) 的整个计算融合在单个 CUDA Kernel 中。这要求我们设计一个像流算法(stream algorithm)那样精心管理片上内存的算法,因为 NVIDIA GPU 的共享内存(shared memory)很小。

对于矩阵乘法这类经典算法,分块用于保证片上内存不超过硬件限制。在 Kernel 执行期间,无论矩阵形状如何,只有 3T2 个元素存储在片上。这种分块方式之所以可行,是因为加法是可结合的(associative),因此整个矩阵乘法可以被分解为许多逐块矩阵乘法之和。

然而,自注意力包含一个不能直接结合的 Softmax 算子,这使得我们无法像矩阵乘法那样简单地分块自注意力。那么,有没有办法让 Softmax 变得可结合呢?

以矩阵乘法 C=A×B 为例,其分块方式如下:矩阵被划分为 T×T 的分块。对每个输出分块,从左到右扫过 A 中相关的分块、从上到下扫过 B 中相关的分块,并把值从全局内存加载到片上内存,整体片上内存占用为 O(T2)。对分块内的部分矩阵乘法,在位置 (i,j) 处,为分块内所有 k 从片上内存加载 A[i,k]B[k,j],然后在片上内存中把 A[i,k]×B[k,j] 累加到 C[i,j]。一个分块计算完成后,把片上的 C 分块写回主内存,再继续处理下一个分块。真实应用中的分块要复杂得多,可参考 CUTLASS 在 A100 上的矩阵乘法实现 [2]。

4.2 分块矩阵乘法的可结合性

我们先从一个非常简单的 GPU 计算开始:矩阵乘法。矩阵乘法要求我们把参与运算的矩阵的行和列反复搬入内存。下面的简单示例展示了如何把第一个矩阵的行搬入内存,以便与第二个矩阵的列相乘。

4.2.1 朴素矩阵乘法

下面的代码是朴素实现,它会反复把同一行或同一列搬入内存,导致很高的缓存未命中率,从而变慢:

python
import numpy as np

def naive_matmul(A, B):
    n, m = A.shape          # A: (n, m)
    m2, p = B.shape         # B: (m, p)
    assert m == m2, "维度不匹配"
    C = np.zeros((n, p))
    for i in range(n):
        for j in range(p):
            for k in range(m):
                C[i, j] += A[i, k] * B[k, j]
    return C

该实现每次累加都重新访问内存,缓存命中率低。序列变长时,这一缺陷会被放大。

4.2.2 分块矩阵乘法

为了提高矩阵乘法的缓存性能,我们使用分块矩阵乘法(tiled matrix multiplication)。这项技术把矩阵分解为能放进高速缓存的小子矩阵(分块,Tile)。我们不顺序计算整个乘积,而是一次处理一个分块,尽可能在较快的内存中复用数据。通过选择合适的分块大小,我们可以显著减少缓存未命中并加快计算。下面的 Python 示例展示了这一方法:

python
import numpy as np

def tiled_matmul(A, B, tile_size):
    n, m = A.shape          # A: (n, m)
    m2, p = B.shape         # B: (m, p)
    assert m == m2, "维度不匹配"

    C = np.zeros((n, p))
    for ii in range(0, n, tile_size):
        for jj in range(0, p, tile_size):
            for kk in range(0, m, tile_size):
                i_end = min(ii + tile_size, n)
                j_end = min(jj + tile_size, p)
                k_end = min(kk + tile_size, m)
                for i in range(ii, i_end):
                    for j in range(jj, j_end):
                        for k in range(kk, k_end):
                            C[i, j] += A[i, k] * B[k, j]
    return C

分块矩阵乘法带来的加速取决于你的硬件和矩阵大小。对于我们的机器学习用例,分块通常能带来显著的加速。因此,当谈到优化矩阵计算——注意力正是其一种特殊形式——关键就在于走向分块矩阵乘法。

那么为什么对注意力很难这样做?虽然我们可以在 QK 之间做分块矩阵乘法,但在最后一次矩阵乘法之前,还需要对结果矩阵做 Softmax。因此,优化注意力的一个关键步骤是弄清楚如何处理这个 Softmax。要开始处理它,我们需要理解这个方程涉及到的复杂性。

4.3 Softmax 的数值稳定性

Softmax 将一组数值转换为概率分布,让较大的数字更突出、较小的数字更不突出。它通过对每个数字取指数来放大差异,然后用所有指数之和去除每个结果,使一切加起来等于 1。

softmax(xi)=exij=1nexj

xi 很大时,exi 会发生上溢(Overflow),产生 。例如 FP16 最大可表示值约为 65,504,而 e11.16.6×104 已超出此范围,导致数值溢出,破坏我们的计算。

安全 Softmax

为了解决这个问题,我们找出张量中的最大值 m,并从指数中减去它。这保证了指数运算不会超出浮点数的表示范围,但确实带来了下溢(underflow)的风险。幸好,因为极小的数字在 Softmax 中会变成 0,这实际上只是一个舍入误差,不会影响我们的计算。

softmax(xi)=eximj=1nexjm

其中 m=maxj=1n(xj)。这样每个 xim0,于是 exim1,不会发生上溢,计算是数值安全的。

为了有效计算 Softmax,我们可以把它写成 3 个步骤。因为需要遍历输入 3 次,所以称之为 3-Pass(三遍)方法:

  1. 求最大值:m=maxj=1n(xj)
  2. 求和:d=j=1nexjm
  3. 归一化:ai=eximd

算法:3-Pass Safe Softmax

符号定义:

符号定义初始值
{mi}i=1N前缀最大值:mi=maxj=1i{xj}m0=
{di}i=1N前缀指数和:di=j=1iexjmNd0=0
{ai}i=1N最终 Softmax 输出:ai=softmax(x)i

dN 是 Safe Softmax 的分母,即 dN=j=1NexjmN

第 1 次遍历:计算前缀最大值(Pass 1)

(2)for i=1N:mimax(mi1, xi)

解释

  • 遍历数组 {x1,x2,,xN}
  • 每一步维护当前最大值 mi
  • 最终得到全局最大值:mN=max1jNxj

第 2 次遍历:累加指数分母(Pass 2)

(3)for i=1N:didi1+eximN

解释

  • 再次遍历数组,利用第 1 步得到的全局最大值 mN
  • 对每个元素计算 eximN 并累加
  • 最终得到 Softmax 的分母:dN=j=1NexjmN

第 3 次遍历:计算最终 Softmax(Pass 3)

(4)for i=1N:aieximNdN

解释

  • 最后一次遍历数组
  • 每个元素计算归一化后的 Softmax 概率
  • 满足:i=1Nai=1,且 i, ai(0,1)

这个算法要求我们对 [1,N] 遍历 3 次。在 Transformer 的自注意力语境中,{xi} 是由 QK 计算出的 Logits。这意味着如果我们不把所有 Logits(即 {xi}i=1N)都存下来(SRAM 不够大装不下全部),就需要访问 QK 三次(即时重算 Logits),这在 I/O 上并不高效。

4.4 Online Softmax:两轮遍历

如果把公式 (2)、(3)、(4) 融合到同一个循环中,就可以把全局内存访问次数从 3 次降到 1 次。遗憾的是,我们不能把公式 (2) 和 (3) 融合进同一个循环,因为公式 (3) 依赖 mN(全局最大值),而 mN 要等第一个循环结束后才能确定。

如果聚焦于第一个循环,我们会发现:我们并不一定需要张量中绝对最大的值,只需要一个足够大、能防止我们溢出的值。通过只找局部最大值、而非全局最大值,我们可以把第 1 遍和第 2 遍融合在一起。

为此,我们可以构造另一个序列 di 作为原始序列 di代理(surrogate)

di:=j=1iexjmi

为什么这样构造?

序列定义特点
原始序列 dij=1iexjmN依赖全局最大值 mN,必须先完整遍历才能得到
代理序列 dij=1iexjmi只依赖当前前缀最大值 mi,可在线递推 ✅

关键性质:两个序列在第 N 项(最后一项)完全相等

dN=dN

因此,公式 (4) 中的 dN 可以安全地替换为 dN

同时我们还能找到 didi1 之间的递推关系:

(5)di=j=1iexjmi(代理序列定义)=(j=1i1exjmi)+eximi(拆出最后一项)=(j=1i1exjmi1emi1mi)+eximi(指数拆分,凑出 di1=(j=1i1exjmi1)=di1emi1mi+eximi(提取公因子)=di1emi1mi+eximi

这个递推形式只依赖 mimi1,因此我们可以在同一个循环中一起计算 midi

算法:2-Pass Online Softmax

初始条件:m0=,d0=0

第 1 次遍历:在线更新最大值与分母(Pass 1)

for i=1N:{mimax(mi1, xi)didi1emi1mi+eximi

解释

公式含义
mi=max(mi1,xi)维护当前前缀最大值 mi(与 3-Pass 算法第 1 步相同)
di=di1emi1mi+eximi在线更新分母核心技巧
由于最大值可能更新,需将旧分母 di1 乘以缩放因子 emi1mi1)进行重标定,再加上新项 eximi

关键性质:遍历结束后,dN=j=1NexjmN,即最终 Softmax 的分母(等价于 3-Pass 算法中的 dN)。

第 2 次遍历:计算最终 Softmax(Pass 2)

for i=1N:aieximNdN

解释

  • 使用第 1 步得到的全局最大值 mN 和分母 dN
  • 每个元素计算归一化的 Softmax 概率值
  • 满足:i=1Nai=1,且 i, ai(0,1)

与 3-Pass Safe Softmax 的对比

特性3-Pass Safe Softmax2-Pass Online Softmax
遍历次数3 次2 次
核心技巧先求 mN 再统一减分母动态重标定:乘以 emi1mi
全局 I/O 次数3N 次数组读写2N 次数组读写
数值稳定性✅ 稳定✅ 同样稳定
数学等价性与标准 Softmax 等价与标准 Softmax 等价

这就是 Online Softmax 论文 [3] 提出的算法:我们边走边算出全局统计量(最大值和分母)。正是这种在线(online)的洞察把我们引向 FlashAttention!

4.5 FlashAttention:单轮遍历

然而它仍然需要两遍才能完成 Softmax 计算——我们能否把遍历次数减少到一遍,以最小化全局 I/O?

回忆一下,注意力可以拆分为 3 步:首先,我们对 QK 做矩阵乘法;然后对结果做 Softmax,得到注意力分数矩阵 A;最后,把注意力矩阵 AV 做矩阵乘法,得到输出矩阵 O

我们已经看到,矩阵乘法可以用分块矩阵乘法加速,Softmax 可以通过 Online Softmax 缩减到 2 遍。在第一遍中,我们对 QK 做矩阵乘法,并确定 X 的最大值以及 Softmax 的分母(基本上就是做 Softmax 的第一遍);第二遍完成 Softmax(构建注意力矩阵 A),并做矩阵乘法得到输出。

现在,如果我们能当场算出 A 的值,就可以合并这两次遍历。虽然 A 当前的表达形式依赖全局统计量,但我们可以使用 Online Softmax 的同一个技巧,把 A 的每个元素改写成局部统计量的结果,然后当我们发现新的最大值时,就按照之前的方式缩放这些值。

这构成了允许两遍合并的新公式:对第 k 行(各行独立),维护运行统计量 midi 与输出累加 oi,并在遍历键序列时用修正因子 emi1mi 在线更新。

算法:Multi-pass Self-Attention(多次遍历自注意力)

符号说明:

符号含义维度/类型
Q[k,:]Q 矩阵的k 行行向量(当前查询 Query)行向量,维度 d
K[:,i]K 矩阵的i 列列向量列向量,维度 d
O[k,:]最终输出 O 矩阵的k 行行向量行向量,维度 d
V[i,:]V 矩阵的i 行行向量(第 i 个 Value)行向量,维度 d
oi部分累加结果向量oi:=j=1iajV[j,:]
表示前 i 个位置 A[k,:]V 的累加和
行向量,维度 d

其中 A[k,:] 是注意力权重矩阵的第 k 行,aj 是其第 j 个元素。

初始条件:

m0=,d0=0,o0=0 (零向量)

第 1 次遍历:计算 xi、前缀最大值 mi、代理分母 di

for i=1N:{xiQ[k,:]K[:,i](点积:Q 的第 k 行 · K 的第 i 行)mimax(mi1, xi)didi1emi1mi+eximi

说明

  • xi = 第 i 个位置的注意力分数(未归一化)
  • 2-Pass Online Softmax 第 1 步完全一致,使用代理分母 di 递推

第 2 次遍历:计算 Softmax 权重并累加 Value 输出(对 i=1N 依次执行以下两步):

(6)aieximNdN(7)oioi1+aiV[i,:]

说明

  • 公式 (6):计算第 i 个位置的 Softmax 注意力权重 ai(标量)
  • 公式 (7):在线累加 Value 向量,无需先完整存储整个注意力矩阵

最终输出:

O[k,:]oN

即第 k 行的最终输出 = 最后一步(i=N)的累加向量 oN

公式推导(FlashAttention 核心思路)

第一步:代入 ai,展开 oi。把公式 (7) 中的 ai 用公式 (6) 的定义替换掉:

(8)oi=j=1i(exjmNdNV[j,:])

问题点:这个形式仍然依赖 mNdN——而它们必须等第一个循环完全结束才能确定,所以仍然需要 2 次遍历。

第二步:再次使用“代理(surrogate)”技巧,构造一个新的代理序列 oi

oi:=j=1i(exjmidiV[j,:])

两个序列的对比:

序列定义依赖
原始 oij=1iexjmNdNV[j,:]依赖全局 mN,dN
代理 oij=1iexjmidiV[j,:]只依赖当前 mi,di

关键性质(与 Softmax 部分类比):遍历到最后时,两者必然相等:

oN=oN=O[k,:]

💡 这就是 FlashAttention 能做到 1-Pass 流式计算的关键突破口!下一步就是推导 oi 的递推式,从而在单次循环里同时更新 mi,di,oi 三个量。

递推关系式推导(公式 (9))

我们需要找到 oioi1 递推的形式:

(9)oi=j=1iexjmidiV[j,:](代理序列定义)=(j=1i1exjmidiV[j,:])前 i1 项+eximidiV[i,:]第 i 项(新加入)=(j=1i1exjmi1di1exjmiexjmi1di1diV[j,:])+eximidiV[i,:](分子分母同乘项凑形式)=(j=1i1exjmi1di1V[j,:])= oi1di1diemi1mi+eximidiV[i,:](化简指数:(xjmi)(xjmi1)=mi1mi=oi1di1emi1midi+eximidiV[i,:]

结论:这个递推式只依赖 di, di1, mi, mi1 以及当前项 xi, V[i,:]——全部可以在单轮循环内获得。因此,我们可以把自注意力的所有计算全部融合到单个循环里!

算法:FlashAttention(1-Pass)

初始条件:

m0=,d0=0,o0=0

单次循环(1-Pass 完成所有计算):

for i=1N:{xiQ[k,:]K[:,i]Q 的第 k 行与 K 的第 i 行做点积)mimax(mi1, xi)(前缀最大值,用于数值稳定)didi1emi1mi+eximi(分母递推)oioi1di1emi1midi+eximidiV[i,:](输出向量递推,公式 (9))

最终输出:

O[k,:]oN

与 Multi-pass Self-Attn 对比

特性Multi-pass Self-AttnFlashAttention(本算法)
遍历次数2 次1 次
全局内存(HBM)访问大量读写 Q,K,V极少 ✅(状态全在 SRAM)
内存占用需存完整 X, A 矩阵O(N) 小状态
分块(Tiling)天然支持
GPU 利用率低(HBM 瓶颈)

4.6 分块版 FlashAttention

实际实现中按块处理,每块包含 b 个 Token,逐 Token 版本可视为 b=1 的特例。对第 k 行(各行独立),设 KBiVBi 分别表示 KV 的第 i 块(第 (i1)b+1ib 行),xi=Q[k,:]KBiRb 为当前块的 Logits。

新增符号说明:

符号中文含义说明
bTile(块)大小每个分块包含的 Token 数量
#tiles一行中分块的总数满足关系:N=b×#tiles
(总序列长度 = 块大小 × 块数量)
xii 个 Tile 的 QK 点积向量存储 Q[k,:]KBi 的计算结果,长度为 b
mi(local)i 个 Tile 内部的局部最大值mi(local)=max1jbxi[j],是 xi 向量内部的最大值

沿用符号mi(全局前缀最大值)、di(代理分母)、oi(代理输出向量)——与 1-Pass FlashAttention 定义一致,仅递推粒度从“单个 Token”升级为“单个 Tile 块”。

初始条件:

m0=,d0=0,o0=0 (零向量)

外层循环:按 Tile(分块)遍历

for i=1#tiles:{xiQ[k,:]KBi(① 取当前 Tile 的 QK 点积向量,长度 bmi(local)=maxj=1b(xi[j])(② 算 Tile 内部局部最大值)mimax(mi1, mi(local))(③ 更新全局最大值)didi1emi1mi+j=1bexi[j]mi(④ 更新全局分母:旧分母缩放 + Tile 内指数和)oioi1di1emi1midi+j=1bexi[j]midiVBi[j,:](⑤ 更新全局输出向量)

第 ⑤ 步逐项解释:

含义
oi1di1emi1midii1 个 Tile 累计的输出向量做重标定(缩放),因为全局最大值和分母可能因新 Tile 加入而改变
j=1bexi[j]midiVBi[j,:]当前 Tile 内部的 b 个 Token做 Softmax + 加权求和(取 VBi 的第 j 行)

最终输出:当处理完所有 Tile 后(i=#tiles=N/b):

O[k,:]oN/b

即第 k 行的最终输出 = 处理完所有分块后的代理输出向量。

Tiling 版本 vs 普通 1-Pass 版本对比

特性普通 1-Pass FlashAttentionFlashAttention(Tiling)(本算法)
处理粒度单个 Token(逐元素)单个 Tile 块(逐 b 个元素)✅
SRAM 占用小,但 HBM 访问仍较多极小Q,K,V 按块读进 SRAM 复用)✅
HBM 全局内存访问较少最少 ✅(FlashAttention 真正的核心优势)
并行友好度一般极高(块间可流水/并行)✅
硬件利用率一般接近峰值

下面的分块参考实现按“逐行 × 按 K/V 分块”组织,与上面的算法一一对应;真实 GPU Kernel 会对 QK/V 同时分块以进一步利用数据复用:

python
import numpy as np

def flash_attention_blocked(Q, K, V, block_size=4):
    """
    分块版 FlashAttention 参考实现(单头)。
    返回 O = softmax(Q @ K.T) @ V,数值上与标准实现一致。
    """
    N, d = K.shape
    O = np.zeros_like(Q)

    for row in range(N):
        q = Q[row, :]
        m_prev = float("-inf")
        d_prev = 0.0
        o_prev = np.zeros(d)

        for start in range(0, N, block_size):
            end = min(start + block_size, N)
            x_block = q @ K[start:end, :].T        # 当前块 Logits,shape (b,)
            m_cur = max(m_prev, x_block.max())     # 更新全局最大值
            p = np.exp(x_block - m_cur)            # 非归一化权重,shape (b,)
            d_cur = d_prev * np.exp(m_prev - m_cur) + p.sum()
            o_cur = (o_prev * d_prev * np.exp(m_prev - m_cur) +
                     p @ V[start:end, :]) / d_cur  # 在线更新输出

            m_prev, d_prev, o_prev = m_cur, d_cur, o_cur

        O[row, :] = o_prev

    return O

FlashAttention 通过在线融合机制,将原本至少三遍的注意力计算压缩到一遍,在 GPU 片上内存中完成计算,避免 O(N2) 的中间显存开销。

5. 技术优势

5.1 显存效率

FlashAttention 避免存储巨大的 N×N 注意力分数矩阵,显著降低显存占用,使在有限显存下处理超长序列(如 100K+ Token 上下文)成为可能。

5.2 计算性能

通过将 QK、Softmax 与 V 乘法融合为单一 CUDA Kernel,FlashAttention 实现了:

  • 大幅减少内存访问次数
  • 计算与内存访问重叠
  • 更高的 GPU 资源利用率

5.3 I/O 复杂度

显存(HBM)访问量往往比算力更早成为瓶颈。标准实现需要将完整的注意力矩阵 A 写入 HBM 再读回,产生 O(N2) 量级的显存访问;FlashAttention 将中间矩阵限制在 SRAM 内,使 HBM 访问量降至 O(N2d2/M),其中 M 为片上内存容量,d 为头维度。在典型配置下(d64128、块大小 64128),计算量(FLOPs)不变,而显存访问量下降一个数量级以上。这正是长序列推理显著提速的根源。

总结

技术遍历次数核心思想
传统 Softmax3 轮分别求最大值、求和、归一化
Online Softmax2 轮在线更新最大值与分母
FlashAttention1 轮(在线融合)直接在线更新输出矩阵,端到端融合

FlashAttention 的创新不仅提升了 Transformer 的训练与推理速度,更为超长上下文理解、文档处理等应用提供了可行的技术方案。

附录:一维向量 Softmax 的推导、证明与实现

设输入为一维向量 x=(x1,x2,,xN),目标是计算:

yi=softmax(x)i=exij=1Nexj

本附录给出 3-Pass(Safe Softmax)、2-Pass(Online Softmax)、1-Pass(延迟归一化)三种算法的推导与正确性证明,并给出只依赖 Python list 的参考实现:除 math.exp 外不调用任何内置聚合函数,最大值与求和均用显式循环完成。

A.1 3-Pass Safe Softmax

推导(平移不变性):对任意常数 c,分子分母同乘 ec,有:

exij=1Nexj=exiecj=1Nexjec=exicj=1Nexjc

即 Softmax 对输入的整体平移保持不变。取 c=m=maxjxj,则 xjm0,于是 exjm(0,1]不会上溢;又因为最大项贡献 emm=1,分母 d1不会除零、也不会整体下溢。这就证明了 Safe Softmax 的数值安全性;由平移不变性,其输出与朴素 Softmax 完全相等。

算法:分三次遍历输入——① 求最大值 m;② 累加指数和 d=j=1Nexjm;③ 归一化 yi=exim/d

python
import math

def softmax_3pass(x):
    """3-Pass Safe Softmax:三次遍历输入向量。"""
    n = len(x)

    # Pass 1:求最大值 m
    m = float("-inf")
    for i in range(n):
        if x[i] > m:
            m = x[i]

    # Pass 2:累加指数和 d,同时缓存指数项
    d = 0.0
    y = [0.0] * n
    for i in range(n):
        y[i] = math.exp(x[i] - m)
        d += y[i]

    # Pass 3:归一化
    for i in range(n):
        y[i] /= d

    return y

A.2 2-Pass Online Softmax

问题:3-Pass 的第 2 次遍历依赖全局最大值 mN,必须等第 1 次遍历结束,两趟无法合并。能否一趟同时算出 mN 和分母?

构造代理序列:定义 di=j=1iexjmi(只依赖前缀最大值 mi),由正文公式 (5) 的推导可得其递推式:

di=di1emi1mi+eximi

于是第 1 趟即可同时在线更新 midi,第 2 趟只做归一化。

正确性证明(数学归纳法):循环不变式为——处理完前 i 个元素后,mi=max1jixjdi=j=1iexjmi

  • 基例i=1):m1=max(,x1)=x1d1=0e+ex1x1=1=j=11exjm1,成立。
  • 归纳步:设不变式对 i1 成立。mi=max(mi1,xi)=max1jixj 显然成立;并且:
di=di1emi1mi+eximi=j=1i1exjmi1emi1mi+eximi=j=1i1exjmi+eximi=j=1iexjmi

i=N 时,dN=j=1NexjmN,正是 Safe Softmax 的分母。因此第 2 趟归一化的输出与 3-Pass 在精确算术下完全一致。

python
import math

def softmax_2pass(x):
    """2-Pass Online Softmax:第 1 趟在线融合最大值与指数和。"""
    n = len(x)

    # Pass 1:在线更新 m_i 与 d_i'
    m = float("-inf")
    d = 0.0
    for i in range(n):
        m_new = x[i] if x[i] > m else m
        d = d * math.exp(m - m_new) + math.exp(x[i] - m_new)
        m = m_new

    # Pass 2:归一化
    y = [0.0] * n
    for i in range(n):
        y[i] = math.exp(x[i] - m) / d

    return y

A.3 1-Pass:延迟归一化

动机:2-Pass 的第 2 次遍历只是为了用 mNdN 做归一化。如果在遍历过程中把每个 exjmi 存进缓冲区,并在最大值更新时对缓冲区整体重标定,那么对输入只需遍历一趟;遍历结束后用最终的 dN 一次性归一化缓冲区即可(归一化只读缓冲区,不再访问输入 x)。

算法:维护运行最大值 m、运行指数和 d、缓冲区 p,保持 p[j]=exjmi;每读入一个 xi,若最大值更新则以因子 emi1mi 重标定整个缓冲区与 d,再追加 eximi

正确性证明(循环不变式归纳):不变式为——处理完前 i 个元素后,对所有 jip[j]=exjmi,且 d=j=1iexjmi

  • 基例i=1):m1=x1p=[ex1x1]=[1]d=1,成立。
  • 归纳步:设不变式对 i1 成立,分两种情形。
    • 情形 1ximi1,最大值不变(mi=mi1),重标定因子 emi1mi=1,缓冲区元素不变,由归纳假设 p[j]=exjmi 保持成立;追加 p[i]=eximid 加上同一项,不变式保持。
    • 情形 2xi>mi1,则 mi=xi。缓冲区中每个 p[j]exjmi1emi1mi=exjmi(指数相加),d 同理;追加 p[i]=eximi=1,不变式保持。
  • 终止i=N):p[j]d=exjmNk=1NexkmN=softmax(x)j,算法正确。

代价分析:输入只读一趟、每个元素只算一次 exp,但最大值每次更新都要 O(i) 重标定缓冲区——最坏情况(输入单调递增)共 O(N2) 次乘法,且需要 O(N) 额外空间。这正是纯 Softmax 无法“免费”单遍的原因。FlashAttention 的巧妙之处(见 4.5 节)在于把归一化融合进输出累加器 oi:累加器大小固定为 d 维、与序列长度无关,重标定代价是 O(d) 而非 O(N),单遍才真正划算。

python
import math

def softmax_1pass(x):
    """1-Pass Online Softmax(延迟归一化):对输入只遍历一趟。"""
    n = len(x)

    m = float("-inf")   # 运行最大值 m_i
    d = 0.0             # 运行指数和 d_i'
    p = []              # 缓冲区:p[j] = e^{x_j - m_i}

    # 唯一一次遍历输入
    for i in range(n):
        m_new = x[i] if x[i] > m else m
        if m_new > m:
            # 最大值更新:重标定缓冲区和指数和
            scale = math.exp(m - m_new)
            for j in range(len(p)):
                p[j] *= scale
            d *= scale
        p.append(math.exp(x[i] - m_new))
        d += p[i]
        m = m_new

    # 用最终分母一次性归一化(只读缓冲区,不再访问 x)
    for i in range(n):
        p[i] /= d

    return p

A.4 运行验证

python
x = [1.0, 3.0, 2.0, 5.0, 4.0, -1.0, 0.5]
print(softmax_3pass(x))
print(softmax_2pass(x))
print(softmax_1pass(x))

big = [1000.0, 1001.0, 1002.0]   # 直接计算 e^1000 必然上溢
print(softmax_1pass(big))

三种实现输出完全一致,且在大值输入下均不发生溢出:

text
[0.0115563, 0.08539015, 0.03141328, 0.63095257, 0.23211448, 0.00156398, 0.00700925]
[0.0115563, 0.08539015, 0.03141328, 0.63095257, 0.23211448, 0.00156398, 0.00700925]
[0.0115563, 0.08539015, 0.03141328, 0.63095257, 0.23211448, 0.00156398, 0.00700925]
[0.09003057, 0.24472847, 0.66524096]

三种算法对比

算法对输入的遍历次数额外空间备注
3-Pass Safe Softmax3O(1)数值安全,但 I/O 最多
2-Pass Online Softmax2O(1)在线融合最大值与分母
1-Pass(延迟归一化)1O(N)最坏 O(N2) 次重标定乘法

可以看到,对纯 Softmax 而言单遍算法要用空间与额外计算换 I/O;而在自注意力中(正文 4.5–4.6 节),由于累加对象是固定维度的输出向量,1-Pass 的在线融合才真正体现出价值。

A.5 1-Pass FlashAttention:一维单行注意力的完整实现

A.3 的延迟归一化说明:纯 Softmax 单遍的代价是缓冲区重标定。把同样的思想用在注意力上——不归一化并保存每个权重 ai,而是把归一化融合进固定维度的输出累加器 oi(正文公式 (9))——就得到了真正实用的 1-Pass 算法。

计算过程(对单个查询向量 q,即输出矩阵的一行):维护运行最大值 m、运行指数和 l(即 di)、输出累加器 o(即 oi),对 KV 只遍历一趟,每个位置依次执行:

  1. xi=qki(点积,求当前位置的 Logit)
  2. mi=max(mi1,xi),重标定因子 s=emi1mi
  3. lls+eximi(分母递推,公式 (5))
  4. oolslnew+eximilnewvi(输出递推,公式 (9))

正确性证明(循环不变式归纳):不变式为——处理完前 i 个键值对后,m=mil=di=j=1iexjmi,且 o=oi=j=1iexjmidivj

  • 基例i=1):s=e=0l=00+ex1x1=1o=00+1v1=ex1x11v1,成立。
  • 归纳步:设不变式对 i1 成立。代码中旧累加器的系数 lslnew 正是公式 (9) 中的 di1emi1midi,新项系数 eximilnew 正是 eximidi;代入公式 (9) 即得更新后的 o=oi,不变式保持。
  • 终止i=N):由 4.5 节的关键性质 oN=oN,输出即 softmax(qK)V 的对应行,算法正确。

与 A.3 的对照:状态只有 (m, l, o),额外空间 O(dv);最大值每次更新的重标定代价是 O(dv) 而非 O(N)——累加器维度固定、与序列长度无关。这正是 FlashAttention 在 GPU 上单遍高效的原因。

Python 实现(纯 list,除 math.exp 外不使用内置聚合函数,点积用显式循环):

python
import math

def dot(a, b):
    """向量点积(显式循环,不用内置聚合函数)。"""
    s = 0.0
    for i in range(len(a)):
        s += a[i] * b[i]
    return s

def flash_attention_1d(q, K, V):
    """
    一维 FlashAttention(1-Pass):单个查询向量 q 对 K/V 做在线注意力。
    q 是长度 d 的 list;K、V 分别是 N 个长度 d、长度 d_v 的 list。
    返回输出向量 o = softmax(qK^T) V 的对应行(长度 d_v)。
    """
    n = len(K)
    dv = len(V[0])

    m = float("-inf")    # 运行最大值 m_i
    l = 0.0              # 运行指数和 d_i'
    o = [0.0] * dv       # 输出累加器 o_i'

    # 对 K/V 只遍历一趟
    for i in range(n):
        x_i = dot(q, K[i])             # ① 当前位置的 Logit
        m_new = x_i if x_i > m else m  # ② 更新前缀最大值
        scale = math.exp(m - m_new)    #    重标定因子(首步为 0)
        p_i = math.exp(x_i - m_new)    #    当前位置的未归一化权重

        l_new = l * scale + p_i        # ③ 在线更新分母(公式 (5))
        w_old = l * scale / l_new      #    旧累加器的重标定系数
        w_new = p_i / l_new            #    当前位置的归一化权重
        for j in range(dv):            # ④ 在线更新输出(公式 (9))
            o[j] = o[j] * w_old + w_new * V[i][j]

        m, l = m_new, l_new

    return o

验证:与朴素的两遍实现(先算全部分数,再归一化加权)对比:

python
def attention_1d_reference(q, K, V):
    """朴素 2-Pass 参考实现:先算全部分数,再做归一化加权。"""
    n = len(K)
    dv = len(V[0])

    # 第一趟:算 Logits,求最大值与分母
    xs = [0.0] * n
    m = float("-inf")
    for i in range(n):
        xs[i] = dot(q, K[i])
        if xs[i] > m:
            m = xs[i]
    l = 0.0
    for i in range(n):
        l += math.exp(xs[i] - m)

    # 第二趟:归一化并加权求和
    o = [0.0] * dv
    for i in range(n):
        a_i = math.exp(xs[i] - m) / l
        for j in range(dv):
            o[j] += a_i * V[i][j]
    return o

q = [0.5, -1.0, 2.0]
K = [[1.0, 0.0, -0.5],
     [0.0, 2.0, 1.0],
     [-1.0, 1.0, 0.0],
     [2.0, -1.0, 1.5]]
V = [[1.0, 2.0],
     [0.5, -1.0],
     [3.0, 0.0],
     [-2.0, 1.0]]

print(flash_attention_1d(q, K, V))
print(attention_1d_reference(q, K, V))

两种实现输出完全一致(另对 n=150、不同 ddv 的随机输入及 x1000 的大值输入做过验证,均一致且不溢出):

text
[-1.96382361, 0.98924009]
[-1.96382361, 0.98924009]

相关链接

参考

  1. Dao, T., Fu, D. Y., Ermon, S., Rudra, A., & Ré, C. (2022). FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. NeurIPS.
  2. NVIDIA. CUTLASS: CUDA Templates for Linear Algebra Subroutines. https://github.com/NVIDIA/cutlass
  3. Milakov, M., & Gimelshein, N. (2018). Online normalizer calculation for softmax. arXiv:1805.02867.
  4. Dao, T. (2023). FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning. arXiv:2307.08691.
  5. Shah, J., Bikshandi, G., Zhang, Y., Thakkar, V., Ramani, P., & Dao, T. (2024). FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision. arXiv:2407.08608.

Maintained by Robin