FlashAttention 详解:从 Online Softmax 到单遍融合
注意力机制是 Transformer 模型的核心组件。所有主流模型架构,如 GPT、LLaMA 和混合专家(Mixture of Experts, MoE),都依赖它来连接 Token 并构建语义。
但注意力是昂贵的。它的计算涉及大规模矩阵乘法,更昂贵的是 GPU 内存与计算单元之间海量的数据移动。随着序列长度不断增长,注意力的标准实现面临显存与计算效率的双重瓶颈。
1. 为什么朴素实现很慢?
回顾注意力运算的定义:
对长度为
朴素实现慢的根源在于 GPU 的内存层级。计算核心只能直接使用一小块非常快的片上内存(SRAM),而更大但更慢的内存池(HBM,即高带宽内存)位于片外。深度学习模型中的大多数操作都是“内存受限”的:它们的速度受限于在 HBM 与 SRAM 之间搬运数据所需的时间,而非在 SRAM 中执行的算术计算。推荐阅读 Horace He 的博客,快速了解深度学习中计算、内存与开销的概念。

注意力的标准实现会在 HBM 中物化
- 从 HBM 读取
和 。 - 计算
,把 写回 HBM。 - 从 HBM 读取
。 - 计算
,把 写回 HBM。 - 从 HBM 读取
和 。 - 计算
,把 写回 HBM。
对于长序列(
问题出在 HBM 流量上。设每个元素占
2. FlashAttention 是什么?
FlashAttention 是一种经过优化的 Transformer 注意力机制,它让注意力在 GPU 上的运行速度大幅提升,同时更加节省内存。

GPU 有两种主要的内存类型。高带宽内存(HBM)容量大但速度相对较慢;片上 SRAM 速度极快但容量非常有限。
标准自注意力在这两者之间不断地搬运数据。这种来回搬运代价高昂,并且随着序列长度的增长,其成本变得十分显著。
FlashAttention 通过将注意力计算拆分为完全能放进快速 SRAM 的小块(Tile)来规避这一问题。每个小块被端到端地完整处理,Softmax 以增量方式应用,因此中间结果无需写回 HBM。由此,完整的注意力矩阵永远不会存储在内存中。
与稀疏注意力或线性注意力方法不同,FlashAttention 不是一种近似方法。它产生的数学输出与标准自注意力完全相同,只是以一种更节省内存的方式执行。
3. FlashAttention 是如何工作的?
FlashAttention 通过重新设计注意力在 GPU 上的计算方式来实现其高效性。它遵循一个简单的机制:在快速片上内存中尽可能多地完成计算,并避免不必要的慢速内存移动。
一个有助于理解的比喻是厨房。GPU 的片上 SRAM 就像一张小而快的厨房台面,是真正准备和烹饪的地方;GPU 高带宽内存(HBM)就像街角的大型杂货店,可以存储你需要的所有东西,但来回往返需要时间。
简单来说,标准注意力每完成一步就要跑到杂货店一趟。相比之下,FlashAttention 会规划好烹饪流程,让你烹饪时所有东西都摆在台面上。

FlashAttention 依赖两个关键思想:分块(Tiling)与重计算(Recomputation)。
3.1 分块(Tiling)
继续用烹饪的例子:分块就是 FlashAttention 把注意力计算塞进小台面的方法。
FlashAttention 不是加载整个序列、构建完整的注意力矩阵,而是把输入拆成小块(Tile)。每个块都能完全放进 GPU 的快速 SRAM 中。FlashAttention 一次处理一个块,从开始到结束,然后才进入下一个块。
回到厨房类比:你无法把一整场宴会的食材都放在小台面上,所以你要小批量地备菜和烹饪——切几样蔬菜、炒熟、清空台面,再处理下一批。通过这种方式,你避免了不断往返杂货店。
这种逐块执行的方式让 FlashAttention 把数据保持在本地,既快又高效,而且永远不会在慢速内存中物化出完整的注意力矩阵。

计算被拆成若干块。矩阵
循环的简化视图:
// 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这种分块结构改变了内存访问模式。算法分块遍历
示例:为了具体化,考虑一块 NVIDIA A100 GPU。它的每个流式多处理器(SM)有 192 KB SRAM。单个 Kernel 会使用其中的一部分,假设工作 SRAM 大小为 M_bytes = 128 KB。如果使用 bfloat16 精度(每个数字 2 字节),有效 SRAM 大小按元素计为 M = 128 * 1024 / 2 = 65,536 个元素。对典型的头维度 d = 64,d = 128,因子为
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)中完成,从而避免
4.1 标准自注意力
忽略注意力头(head)与批次(batch)维度(这两个维度上的计算完全并行),同时省略注意力掩码(mask)和缩放因子
其中
标准的计算方式是把自注意力分解为三步:
其中
直接实现需要把完整的
FlashAttention 的关键在于:它无需在全局内存中物化
对于矩阵乘法这类经典算法,分块用于保证片上内存不超过硬件限制。在 Kernel 执行期间,无论矩阵形状如何,只有
然而,自注意力包含一个不能直接结合的 Softmax 算子,这使得我们无法像矩阵乘法那样简单地分块自注意力。那么,有没有办法让 Softmax 变得可结合呢?
以矩阵乘法
4.2 分块矩阵乘法的可结合性
我们先从一个非常简单的 GPU 计算开始:矩阵乘法。矩阵乘法要求我们把参与运算的矩阵的行和列反复搬入内存。下面的简单示例展示了如何把第一个矩阵的行搬入内存,以便与第二个矩阵的列相乘。
4.2.1 朴素矩阵乘法
下面的代码是朴素实现,它会反复把同一行或同一列搬入内存,导致很高的缓存未命中率,从而变慢:
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 示例展示了这一方法:
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分块矩阵乘法带来的加速取决于你的硬件和矩阵大小。对于我们的机器学习用例,分块通常能带来显著的加速。因此,当谈到优化矩阵计算——注意力正是其一种特殊形式——关键就在于走向分块矩阵乘法。
那么为什么对注意力很难这样做?虽然我们可以在
4.3 Softmax 的数值稳定性
Softmax 将一组数值转换为概率分布,让较大的数字更突出、较小的数字更不突出。它通过对每个数字取指数来放大差异,然后用所有指数之和去除每个结果,使一切加起来等于 1。
当
安全 Softmax
为了解决这个问题,我们找出张量中的最大值
其中
为了有效计算 Softmax,我们可以把它写成 3 个步骤。因为需要遍历输入 3 次,所以称之为 3-Pass(三遍)方法:
- 求最大值:
- 求和:
- 归一化:
算法:3-Pass Safe Softmax
符号定义:
| 符号 | 定义 | 初始值 |
|---|---|---|
| 前缀最大值: | ||
| 前缀指数和: | ||
| 最终 Softmax 输出: | — |
注:
是 Safe Softmax 的分母,即
第 1 次遍历:计算前缀最大值(Pass 1)
解释:
- 遍历数组
- 每一步维护当前最大值
- 最终得到全局最大值:
第 2 次遍历:累加指数分母(Pass 2)
解释:
- 再次遍历数组,利用第 1 步得到的全局最大值
- 对每个元素计算
并累加 - 最终得到 Softmax 的分母:
第 3 次遍历:计算最终 Softmax(Pass 3)
解释:
- 最后一次遍历数组
- 每个元素计算归一化后的 Softmax 概率
- 满足:
,且
这个算法要求我们对
4.4 Online Softmax:两轮遍历
如果把公式 (2)、(3)、(4) 融合到同一个循环中,就可以把全局内存访问次数从 3 次降到 1 次。遗憾的是,我们不能把公式 (2) 和 (3) 融合进同一个循环,因为公式 (3) 依赖
如果聚焦于第一个循环,我们会发现:我们并不一定需要张量中绝对最大的值,只需要一个足够大、能防止我们溢出的值。通过只找局部最大值、而非全局最大值,我们可以把第 1 遍和第 2 遍融合在一起。
为此,我们可以构造另一个序列
为什么这样构造?
| 序列 | 定义 | 特点 |
|---|---|---|
| 原始序列 | 依赖全局最大值 | |
| 代理序列 | 只依赖当前前缀最大值 |
关键性质:两个序列在第
因此,公式 (4) 中的
可以安全地替换为 !
同时我们还能找到
这个递推形式只依赖
算法:2-Pass Online Softmax
初始条件:
第 1 次遍历:在线更新最大值与分母(Pass 1)
解释:
| 公式 | 含义 |
|---|---|
| 维护当前前缀最大值 | |
| 在线更新分母核心技巧: 由于最大值可能更新,需将旧分母 |
关键性质:遍历结束后,
,即最终 Softmax 的分母(等价于 3-Pass 算法中的 )。
第 2 次遍历:计算最终 Softmax(Pass 2)
解释:
- 使用第 1 步得到的全局最大值
和分母 - 每个元素计算归一化的 Softmax 概率值
- 满足:
,且
与 3-Pass Safe Softmax 的对比:
| 特性 | 3-Pass Safe Softmax | 2-Pass Online Softmax |
|---|---|---|
| 遍历次数 | 3 次 | 2 次 ✅ |
| 核心技巧 | 先求 | 分母动态重标定:乘以 |
| 全局 I/O 次数 | ||
| 数值稳定性 | ✅ 稳定 | ✅ 同样稳定 |
| 数学等价性 | 与标准 Softmax 等价 | 与标准 Softmax 等价 |
这就是 Online Softmax 论文 [3] 提出的算法:我们边走边算出全局统计量(最大值和分母)。正是这种在线(online)的洞察把我们引向 FlashAttention!
4.5 FlashAttention:单轮遍历
然而它仍然需要两遍才能完成 Softmax 计算——我们能否把遍历次数减少到一遍,以最小化全局 I/O?
回忆一下,注意力可以拆分为 3 步:首先,我们对
我们已经看到,矩阵乘法可以用分块矩阵乘法加速,Softmax 可以通过 Online Softmax 缩减到 2 遍。在第一遍中,我们对
现在,如果我们能当场算出
这构成了允许两遍合并的新公式:对第
算法:Multi-pass Self-Attention(多次遍历自注意力)
符号说明:
| 符号 | 含义 | 维度/类型 |
|---|---|---|
| 行向量,维度 | ||
| 列向量,维度 | ||
| 最终输出 | 行向量,维度 | |
| 行向量,维度 | ||
| 部分累加结果向量: 表示前 | 行向量,维度 |
其中
是注意力权重矩阵的第 行, 是其第 个元素。
初始条件:
第 1 次遍历:计算
说明:
= 第 个位置的注意力分数(未归一化) - 与 2-Pass Online Softmax 第 1 步完全一致,使用代理分母
递推
第 2 次遍历:计算 Softmax 权重并累加 Value 输出(对
说明:
- 公式 (6):计算第
个位置的 Softmax 注意力权重 (标量) - 公式 (7):在线累加 Value 向量,无需先完整存储整个注意力矩阵
最终输出:
即第
公式推导(FlashAttention 核心思路)
第一步:代入
问题点:这个形式仍然依赖
第二步:再次使用“代理(surrogate)”技巧,构造一个新的代理序列
两个序列的对比:
| 序列 | 定义 | 依赖 |
|---|---|---|
| 原始 | 依赖全局 | |
| 代理 | 只依赖当前 |
关键性质(与 Softmax 部分类比):遍历到最后时,两者必然相等:
💡 这就是 FlashAttention 能做到 1-Pass 流式计算的关键突破口!下一步就是推导
的递推式,从而在单次循环里同时更新 三个量。
递推关系式推导(公式 (9))
我们需要找到
结论:这个递推式只依赖
算法:FlashAttention(1-Pass)
初始条件:
单次循环(1-Pass 完成所有计算):
最终输出:
与 Multi-pass Self-Attn 对比:
| 特性 | Multi-pass Self-Attn | FlashAttention(本算法) |
|---|---|---|
| 遍历次数 | 2 次 | 1 次 ✅ |
| 全局内存(HBM)访问 | 大量读写 | 极少 ✅(状态全在 SRAM) |
| 内存占用 | 需存完整 | |
| 分块(Tiling) | 难 | 天然支持 ✅ |
| GPU 利用率 | 低(HBM 瓶颈) | 高 ✅ |
4.6 分块版 FlashAttention
实际实现中按块处理,每块包含
新增符号说明:
| 符号 | 中文含义 | 说明 |
|---|---|---|
| Tile(块)大小 | 每个分块包含的 Token 数量 | |
| 一行中分块的总数 | 满足关系: (总序列长度 = 块大小 × 块数量) | |
| 第 | 存储 | |
| 第 |
沿用符号:
(全局前缀最大值)、 (代理分母)、 (代理输出向量)——与 1-Pass FlashAttention 定义一致,仅递推粒度从“单个 Token”升级为“单个 Tile 块”。
初始条件:
外层循环:按 Tile(分块)遍历
第 ⑤ 步逐项解释:
| 项 | 含义 |
|---|---|
| 对前 | |
| 对当前 Tile 内部的 |
最终输出:当处理完所有 Tile 后(
即第
Tiling 版本 vs 普通 1-Pass 版本对比:
| 特性 | 普通 1-Pass FlashAttention | FlashAttention(Tiling)(本算法) |
|---|---|---|
| 处理粒度 | 单个 Token(逐元素) | 单个 Tile 块(逐 |
| SRAM 占用 | 小,但 HBM 访问仍较多 | 极小( |
| HBM 全局内存访问 | 较少 | 最少 ✅(FlashAttention 真正的核心优势) |
| 并行友好度 | 一般 | 极高(块间可流水/并行)✅ |
| 硬件利用率 | 一般 | 接近峰值 ✅ |
下面的分块参考实现按“逐行 × 按
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 OFlashAttention 通过在线融合机制,将原本至少三遍的注意力计算压缩到一遍,在 GPU 片上内存中完成计算,避免
5. 技术优势
5.1 显存效率
FlashAttention 避免存储巨大的
5.2 计算性能
通过将
- 大幅减少内存访问次数
- 计算与内存访问重叠
- 更高的 GPU 资源利用率
5.3 I/O 复杂度
显存(HBM)访问量往往比算力更早成为瓶颈。标准实现需要将完整的注意力矩阵
总结
| 技术 | 遍历次数 | 核心思想 |
|---|---|---|
| 传统 Softmax | 3 轮 | 分别求最大值、求和、归一化 |
| Online Softmax | 2 轮 | 在线更新最大值与分母 |
| FlashAttention | 1 轮(在线融合) | 直接在线更新输出矩阵,端到端融合 |
FlashAttention 的创新不仅提升了 Transformer 的训练与推理速度,更为超长上下文理解、文档处理等应用提供了可行的技术方案。
附录:一维向量 Softmax 的推导、证明与实现
设输入为一维向量
本附录给出 3-Pass(Safe Softmax)、2-Pass(Online Softmax)、1-Pass(延迟归一化)三种算法的推导与正确性证明,并给出只依赖 Python list 的参考实现:除 math.exp 外不调用任何内置聚合函数,最大值与求和均用显式循环完成。
A.1 3-Pass Safe Softmax
推导(平移不变性):对任意常数
即 Softmax 对输入的整体平移保持不变。取
算法:分三次遍历输入——① 求最大值
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 yA.2 2-Pass Online Softmax
问题:3-Pass 的第 2 次遍历依赖全局最大值
构造代理序列:定义
于是第 1 趟即可同时在线更新
正确性证明(数学归纳法):循环不变式为——处理完前
- 基例(
): , ,成立。 - 归纳步:设不变式对
成立。 显然成立;并且:
当
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 yA.3 1-Pass:延迟归一化
动机:2-Pass 的第 2 次遍历只是为了用
算法:维护运行最大值
正确性证明(循环不变式归纳):不变式为——处理完前
- 基例(
): , , ,成立。 - 归纳步:设不变式对
成立,分两种情形。 - 情形 1:
,最大值不变( ),重标定因子 ,缓冲区元素不变,由归纳假设 保持成立;追加 , 加上同一项,不变式保持。 - 情形 2:
,则 。缓冲区中每个 (指数相加), 同理;追加 ,不变式保持。
- 情形 1:
- 终止(
): ,算法正确。
代价分析:输入只读一趟、每个元素只算一次 exp,但最大值每次更新都要
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 pA.4 运行验证
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))三种实现输出完全一致,且在大值输入下均不发生溢出:
[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 Softmax | 3 | 数值安全,但 I/O 最多 | |
| 2-Pass Online Softmax | 2 | 在线融合最大值与分母 | |
| 1-Pass(延迟归一化) | 1 | 最坏 |
可以看到,对纯 Softmax 而言单遍算法要用空间与额外计算换 I/O;而在自注意力中(正文 4.5–4.6 节),由于累加对象是固定维度的输出向量,1-Pass 的在线融合才真正体现出价值。
A.5 1-Pass FlashAttention:一维单行注意力的完整实现
A.3 的延迟归一化说明:纯 Softmax 单遍的代价是缓冲区重标定。把同样的思想用在注意力上——不归一化并保存每个权重
计算过程(对单个查询向量
(点积,求当前位置的 Logit) ,重标定因子 (分母递推,公式 (5)) (输出递推,公式 (9))
正确性证明(循环不变式归纳):不变式为——处理完前
- 基例(
): , , ,成立。 - 归纳步:设不变式对
成立。代码中旧累加器的系数 正是公式 (9) 中的 ,新项系数 正是 ;代入公式 (9) 即得更新后的 ,不变式保持。 - 终止(
):由 4.5 节的关键性质 ,输出即 的对应行,算法正确。
与 A.3 的对照:状态只有
Python 实现(纯 list,除 math.exp 外不使用内置聚合函数,点积用显式循环):
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验证:与朴素的两遍实现(先算全部分数,再归一化加权)对比:
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))两种实现输出完全一致(另对
[-1.96382361, 0.98924009]
[-1.96382361, 0.98924009]相关链接
参考
- Dao, T., Fu, D. Y., Ermon, S., Rudra, A., & Ré, C. (2022). FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. NeurIPS.
- NVIDIA. CUTLASS: CUDA Templates for Linear Algebra Subroutines. https://github.com/NVIDIA/cutlass
- Milakov, M., & Gimelshein, N. (2018). Online normalizer calculation for softmax. arXiv:1805.02867.
- Dao, T. (2023). FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning. arXiv:2307.08691.
- 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.