1. FlashAttention 详解(零基础入门)
    1. 1. 1. 先搞清楚这些基础概念
      1. 1.1. 1.1 大语言模型里的”Attention”是什么
      2. 1.2. 1.2 所有缩写和术语一次说清楚
    2. 2. 2. Attention 的完整流程(看懂每个公式)
      1. 2.1. 2.1 输入:一段文本变成了什么
      2. 2.2. 2.2 第一步:生成 Q、K、V 三个矩阵
      3. 2.3. 2.3 第二步:计算注意力分数 S
      4. 2.4. 2.4 第三步:用 softmax 把分数变成概率
      5. 2.5. 2.5 第四步:用权重汇总 V 得到输出
      6. 2.6. 2.6 完整公式总览
    3. 3. 3. 标准 Attention 的两大痛点
      1. 3.1. 3.1 痛点一:N×N 矩阵吃显存
      2. 3.2. 3.2 痛点二:HBM 带宽成为瓶颈
    4. 4. 4. FlashAttention 的核心洞察
      1. 4.1. 4.1 关键观察
      2. 4.2. 4.2 拦路虎:softmax 需要”全行信息”
    5. 5. 5. Online Softmax 详解
      1. 5.1. 5.1 核心思想
      2. 5.2. 5.2 数学推导
      3. 5.3. 5.3 一个具体数值例子
      4. 5.4. 5.4 顺便更新输出 O
    6. 6. 6. FlashAttention 完整算法
      1. 6.1. 6.1 核心思路
      2. 6.2. 6.2 伪代码
      3. 6.3. 6.3 关键优势
    7. 7. 7. 效果对比
      1. 7.1. 7.1 数字上的对比(N=8192,d=128)
      2. 7.2. 7.2 为什么快这么多
    8. 8. 8. 版本演进
    9. 9. 9. 一句话总结
    10. 10. 10. 延伸阅读

FlashAttention 详解(零基础入门)

本文由 AI 生成,内容可能存在错误或不准确之处,请结合可靠来源核验后再作参考。

NOTE

阅读目标:完全不懂大模型的人,看完能搞清楚:

  • Attention(注意力机制)是什么、怎么算
  • 为什么标准 Attention 在长文本上又慢又耗显存
  • FlashAttention 做了什么改动、为什么能快几倍
  • 每一个公式、每一个缩写都解释清楚

1. 先搞清楚这些基础概念

1.1 大语言模型里的”Attention”是什么

想象你在读一个句子:“小明把书放在桌上,他走了。”

当你读到”“时,你的大脑会自动”看回去”——“他”指的是小明。这个”看回去、找相关信息”的动作,就叫注意力(Attention)

大语言模型处理文本的每一个词时,都要做类似的事:

  • 看当前这个词
  • 扫视前面所有词
  • 决定”要给前面每个词分配多少注意力”
  • 用分配好的权重,把前面词的信息加权汇总

Attention 机制就是这个”扫视 + 加权汇总”的数学公式实现。

1.2 所有缩写和术语一次说清楚

缩写 / 符号 英文全称 中文含义 说明
LLM Large Language Model 大语言模型 GPT、Llama、Qwen 都是
Transformer - 变换器(模型架构名) 现代 LLM 的基础架构,核心就是 Attention
Token - 词元 文本被切成的最小单位,一个汉字或一个英文单词通常是 1-2 个 token
N sequence length 序列长度 一句话有多少个 token
d dimension 维度 每个 token 用多长的向量表示(常见 64 / 128)
Q Query 查询 “我想问什么”
K Key “我能被问什么”
V Value “我实际提供什么信息”
S Score 注意力分数 Q 和 K 匹配度,越高越相关
P Probability 注意力权重 把 S 归一化成 0-1 的概率
O Output 输出 最终加权汇总的结果
MatMul Matrix Multiplication 矩阵乘法
softmax - 归一化指数函数 把一堆数变成加起来等于 1 的概率
GPU Graphics Processing Unit 图形处理器 大模型主要在 GPU 上跑
HBM High Bandwidth Memory 高带宽内存 GPU 的”大仓库”内存,容量大但慢(相对)
SRAM Static RAM 静态随机存储 GPU 里的”小厨房”内存,容量小但极快
FP16 Floating Point 16-bit 16 位浮点数 每个数占 2 字节
FP32 Floating Point 32-bit 32 位浮点数 每个数占 4 字节
FLOPs Floating Point Operations 浮点运算次数 衡量计算量的单位

2. Attention 的完整流程(看懂每个公式)

2.1 输入:一段文本变成了什么

假设句子有 $N=4$ 个 token,每个 token 用 $d=3$ 维向量表示。输入矩阵为:

$$
X=
\begin{bmatrix}
x_{11} & x_{12} & x_{13} \\
x_{21} & x_{22} & x_{23} \\
x_{31} & x_{32} & x_{33} \\
x_{41} & x_{42} & x_{43}
\end{bmatrix}
\in \mathbb{R}^{N\times d}
$$

  • 每一代表一个 token 的向量表示
  • 共 4 行(4 个 token),每行 3 维
  • 数学符号写作 $X\in\mathbb{R}^{N\times d}$,读作“$X$ 是 $N$ 行 $d$ 列的实数矩阵”

2.2 第一步:生成 Q、K、V 三个矩阵

每个 token 会被映射成三个不同的向量:Query、Key、Value。

$$
\begin{aligned}
Q &= XW_Q, \\
K &= XW_K, \\
V &= XW_V.
\end{aligned}
$$

逐个解释

  • $X$:输入矩阵,形状为 $[N,d]$,每行对应一个 token
  • $W_Q,W_K,W_V$:三个可学习的权重矩阵(训练时学出来的),形状均为 $[d,d]$
  • $Q,K,V$:运算结果,形状均为 $[N,d]$

形象理解

  • Q(Query):像每个 token 伸出的”触角”,代表”我在找什么”
  • K(Key):像每个 token 贴的”标签”,代表”我是什么”
  • V(Value):代表”如果你关注我,你能得到什么信息”

2.3 第二步:计算注意力分数 S

Q 去和所有 K 做匹配:

$$
S=\frac{QK^\mathsf{T}}{\sqrt{d}}.
$$

逐个解释

  • $Q$:形状为 $[N,d]$
  • $K^\mathsf{T}$:$K$ 的转置(行列互换),形状为 $[d,N]$
  • $QK^\mathsf{T}$:矩阵乘法(MatMul),结果形状为 $[N,N]$
  • $\sqrt{d}$:缩放因子,防止数值过大导致后续 softmax 饱和
  • $S$:注意力分数矩阵,形状为 $[N,N]$

$S$ 里的每个元素 $S_{ij}$ 表示什么?

$$
S_{ij}=\frac{Q_i\cdot K_j}{\sqrt{d}},
$$

它表示“第 $i$ 个 token 对第 $j$ 个 token 的关注度(尚未归一化)”。

举例($N=4$):

token 1 token 2 token 3 token 4
token 1 对…的分数 0.8 0.2 0.5 0.1
token 2 对…的分数 0.3 0.9 0.2 0.4
token 3 对…的分数 0.1 0.4 0.7 0.6
token 4 对…的分数 0.2 0.3 0.5 0.8

整张表就是 $S$ 矩阵,是一个 $N\times N$(这里为 $4\times4$)的方阵。

2.4 第三步:用 softmax 把分数变成概率

$$
P=\operatorname{softmax}(S).
$$

softmax 的作用:把一行数字变成加起来等于 1 的概率分布。

对每一行 i,softmax 公式是:

$$
P_{ij}=\frac{\exp(S_{ij})}{\sum_{k=1}^{N}\exp(S_{ik})}.
$$

逐个符号解释

  • $\exp(x)$:指数函数 $e^x$,其中 $e\approx2.718$
  • 分子:当前元素的指数
  • 分母:这一行所有元素指数的总和
  • $\sum$:连加求和符号,$\sum_{k=1}^{N}$ 表示“对 $k$ 从 $1$ 累加到 $N$”
  • $P_{ij}$:归一化后的概率,满足每行元素之和为 $1$

数值稳定性问题(重要!):直接计算 $\exp$ 可能溢出(例如 $\exp(1000)$ 是天文数字),实际使用的公式为:

$$
P_{ij}=\frac{\exp(S_{ij}-m_i)}{\sum_{k=1}^{N}\exp(S_{ik}-m_i)}.
$$

其中 $m_i=\max_k S_{ik}$ 是第 $i$ 行的最大值。

减去最大值后结果不变(数学上恒等),但可以保证 $\exp$ 的输入值不大于 $0$,从而避免溢出。这点后面讲 FlashAttention 时会很关键

2.5 第四步:用权重汇总 V 得到输出

$$
O=PV.
$$

  • $P$:形状为 $[N,N]$,表示注意力权重
  • $V$:形状为 $[N,d]$,表示值矩阵
  • $O$:形状为 $[N,d]$,表示最终输出

$O$ 的每一行 $O_i$ 是什么?

$$
O_i=\sum_{j=1}^{N}P_{ij}V_j.
$$

它表示“第 $i$ 个 token 的输出等于所有 token 的 $V$ 按照注意力权重进行加权求和”。

2.6 完整公式总览

$$
\operatorname{Attention}(Q,K,V)
=\operatorname{softmax}\left(\frac{QK^\mathsf{T}}{\sqrt{d}}\right)V.
$$

TIP

整个过程可以用三步概括

  1. 打分:$S = QK^\mathsf{T}/\sqrt{d}$,得到 $N\times N$ 的分数矩阵
  2. 归一化:$P = \operatorname{softmax}(S)$,每行元素之和为 $1$
  3. 汇总:$O = PV$,按权重对 $V$ 加权求和

3. 标准 Attention 的两大痛点

3.1 痛点一:N×N 矩阵吃显存

$N$ 是序列长度。下面列出 $S$ 和 $P$ 这两个中间矩阵的大小:

序列长度 N S 矩阵大小(FP16) 说明
512 0.5 MB 尚可
2,048 8 MB OK
4,096 32 MB 开始紧张
8,192 128 MB 紧张
16,384 512 MB 非常紧张
32,768 2 GB 爆炸
131,072 32 GB 直接 OOM

计算方式为 $N\times N\times 2$ 字节(FP16)。

每个 Attention 层都要这么一个矩阵,而一个 LLM 有几十层——总显存需求随 N 平方增长,这就是为什么早期 Transformer 根本跑不了长文本。

3.2 痛点二:HBM 带宽成为瓶颈

GPU 的内存有两种:

类型 位置 容量 速度
HBM 显存条 几十 GB ~ 1-2 TB/s
SRAM 在芯片上 几十 MB(超小) ~ 10-20 TB/s

HBM 虽然叫”高带宽内存”,但比 SRAM 还是慢 10 倍。

标准 Attention 在 HBM 和 SRAM 之间来回搬数据:

  1. 计算 $S=QK^\mathsf{T}/\sqrt{d}$,再将 $N\times N$ 的 $S$ 写入 HBM。
  2. 从 HBM 读取 $S$,计算 $P=\operatorname{softmax}(S)$,再将 $P$ 写回 HBM。
  3. 从 HBM 读取 $P$,计算 $O=PV$,最后写回 $O$。

每一步都要读写 $N\times N$ 的大矩阵,HBM 带宽被压满。而此时 GPU 的计算单元(MAC 阵列)大部分时间在等数据,算力利用率很低。

IMPORTANT

核心问题:Attention 不是“算得慢”,而是“数据搬得慢”。典型场景下 GPU 计算单元利用率只有 10–30%。


4. FlashAttention 的核心洞察

4.1 关键观察

IMPORTANT

我们最终只需要 $O = PV$ 这个形状为 $[N,d]$ 的结果。中间形状为 $[N,N]$ 的 $S$ 和 $P$ 矩阵,其实不需要完整存在。

如果我们能一边算一点 S、一边做 softmax、一边和对应的 V 相乘,中间结果全部待在 SRAM 里,永不写回 HBM——那 HBM 流量就能大幅降低。

4.2 拦路虎:softmax 需要”全行信息”

问题来了:

$$
P_{ij}=\frac{\exp(S_{ij}-m_i)}{\sum_{k=1}^{N}\exp(S_{ik}-m_i)}.
$$

要算这个,你得知道:

  • $m_i$:整行 $S$ 的最大值
  • 分母:整行 S 的指数之和

你必须看到整行才能算 softmax。如果只有前半行,算出来的 softmax 是错的。

这个约束让朴素的”分块算法”行不通——FlashAttention 的核心贡献就是解决了这个问题,用的是 Online Softmax 技巧。


5. Online Softmax 详解

5.1 核心思想

边看数据边维护”目前为止的 max 和 sum”,新数据来的时候修正之前的结果

5.2 数学推导

假设一整行有 $N$ 个数,我们分两块来看:前一块记作 $x^{(1)}$,后一块记作 $x^{(2)}$。

处理完第一块后,我们维护两个统计量:

$$
\begin{aligned}
m^{(1)} &= \max_j x_j^{(1)}, \\
\ell^{(1)} &= \sum_j \exp\left(x_j^{(1)}-m^{(1)}\right).
\end{aligned}
$$

第二块到来时,我们想要的新统计量是 “整个(第一块+第二块)的 max 和 sum”

新的 max:

$$
m^{(2)}=\max\left(m^{(1)},\max_j x_j^{(2)}\right).
$$

新的 sum(这里是关键):

$$
\ell^{(2)}
=\underset{\text{修正旧的 sum}}{\underbrace{\ell^{(1)}\exp\left(m^{(1)}-m^{(2)}\right)}}
+\underset{\text{加入新块贡献}}{\underbrace{\sum_{j}\exp\left(x_{j}^{(2)}-m^{(2)}\right)}}.
$$

为什么是这个公式? 因为旧的 $\ell^{(1)}$ 是以旧的最大值为基准计算的(即每项减去 $m^{(1)}$),现在需要切换到以新最大值为基准。对单项做如下变换:

$$
\begin{aligned}
\exp\left(x_j-m^{(1)}\right)
&=\exp\left(x_j-m^{(2)}+m^{(2)}-m^{(1)}\right) \\
&=\exp\left(x_j-m^{(2)}\right)\exp\left(m^{(2)}-m^{(1)}\right).
\end{aligned}
$$

因此,旧的 sum 乘以 $\exp\left(m^{(1)}-m^{(2)}\right)$,就能修正到新的基准。

WARNING

这个转换是精确的,不是近似;其数值结果与一次性处理全部数据完全相同。

5.3 一个具体数值例子

假设一行共有 $4$ 个数 $[1,3,2,4]$,分两块处理,每块包含 $2$ 个数。

一次性算(参考答案)

  • $\max(x)=4$
  • 减去最大值后得到 $[-3,-1,-2,0]$
  • 取指数后得到 $[0.0498,0.3679,0.1353,1.0]$
  • 指数之和约为 $1.553$
  • softmax 结果约为 $[0.032,0.237,0.087,0.644]$

Online 算

第一块 $[1,3]$:

  • $m^{(1)}=3$
  • $\ell^{(1)}=\exp(1-3)+\exp(3-3)=0.1353+1.0=1.1353$

第二块 $[2,4]$:

  • 新的最大值为 $m^{(2)}=\max(3,4)=4$
  • 修正旧的 sum:$1.1353\times\exp(3-4)=1.1353\times0.3679=0.4178$
  • 新块贡献:$\exp(2-4)+\exp(4-4)=0.1353+1.0=1.1353$
  • $\ell^{(2)}=0.4178+1.1353=1.5531$

1.5531 ≈ 1.553,和一次性算的结果完全一致,说明 online 算法是精确的。

5.4 顺便更新输出 O

Attention 不止算 softmax,还要乘以 V。Online 地更新 O 的公式:

设当前已经处理到第 $i$ 个 tile,并维护累积量 $O^{(i)}$:

$$
O^{(i)}
=\frac{
O^{(i-1)}\ell^{(i-1)}\exp\left(m^{(i-1)}-m^{(i)}\right)
+\exp\left(S^{(i)}-m^{(i)}\right)V^{(i)}
}{\ell^{(i)}}.
$$

看着复杂,本质就三件事:

  1. 旧的 $O$ 重新缩放,使用与更新 $\ell$ 相同的 $\exp\left(m^{(i-1)}-m^{(i)}\right)$ 修正因子
  2. 加上当前 tile 的贡献 $\exp\left(S^{(i)}-m^{(i)}\right)V^{(i)}$
  3. 使用新的 $\ell$ 归一化

6. FlashAttention 完整算法

6.1 核心思路

  • Q 按行切块(外层循环)
  • K, V 按行切块(内层循环)
  • 每个 Q 块在 SRAM 里常驻,扫一遍所有 K/V 块
  • 维护 m, ℓ, O 三个累积量,online 更新
  • S 和 P 矩阵永远不完整出现,只有小 tile 大小

6.2 伪代码

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
## 输入:Q, K, V ∈ [N, d],都在 HBM
## 输出:O ∈ [N, d]

## 分块大小(能装进 SRAM 即可)
B_q = 64 # Q 每块行数
B_k = 64 # K/V 每块行数

for i in range(N // B_q): # 外层:Q 块
Q_i = 从 HBM 读取 Q 的第 i 块 [B_q, d] # 一次 HBM 读
O_i = zeros(B_q, d) # SRAM 里初始化
m_i = -infinity # 当前最大值
l_i = 0 # 当前 sum

for j in range(N // B_k): # 内层:K, V 块
K_j = 从 HBM 读取 K 的第 j 块 [B_k, d]
V_j = 从 HBM 读取 V 的第 j 块 [B_k, d]

# 在 SRAM 里完成所有计算
S_ij = Q_i @ K_j.T / sqrt(d) # [B_q, B_k]
m_ij = max(S_ij, axis=-1) # 当前块的 max

m_new = max(m_i, m_ij) # 新 max
P_ij = exp(S_ij - m_new) # 局部 softmax 分子
l_new = l_i * exp(m_i - m_new) + sum(P_ij)

# 更新 O
O_i = O_i * (l_i / l_new) * exp(m_i - m_new) + (P_ij @ V_j) / l_new

m_i, l_i = m_new, l_new

把 O_i 写回 HBM # 一次 HBM 写

6.3 关键优势

  • $S$ 和 $P$ 从未以 $[N,N]$ 的完整形态存在,只有 $[B_q,B_k]$ 的小 tile
  • HBM 读写复杂度:$O(Nd)$(只读写 $Q$、$K$、$V$、$O$),不再是 $O(N^2)$
  • SRAM 充分利用:所有中间计算都在快内存里

7. 效果对比

7.1 数字上的对比(N=8192,d=128)

指标 原始 Attention FlashAttention 改善
S 矩阵显存占用 128 MB / head 0(不存)
HBM 读写量 ~ 540 MB ~ 35 MB 15×
实际速度(A100 GPU) 基准 2-4× 快 快 3 倍
可用最大上下文 2K-8K 32K+

7.2 为什么快这么多

TIP

Attention 本来是“memory-bound”(内存带宽瓶颈),不是“compute-bound”(算力瓶颈)。
FlashAttention 没增加计算量,甚至略微增加了一点点(重复计算 exp),但大幅减少了 HBM 流量——于是带宽瓶颈解除,GPU 的计算单元终于跑满了。


8. 版本演进

版本 年份 核心改进
FlashAttention-1 2022 首创 tiling + online softmax,比原生 attention 快 2-4×
FlashAttention-2 2023 换了循环顺序(Q 在外层),warp 级并行更好,比 v1 再快 2×
FlashAttention-3 2024 针对 H100 异步 GEMM + FP8 支持,又快 1.5-2×

三个版本算法骨架都一样,区别在于针对不同硬件(A100 / H100)的微架构做了精细调优。


9. 一句话总结

IMPORTANT

FlashAttention 的核心贡献
证明了“Attention 的 $N\times N$ 矩阵不需要物化”——通过分块计算 + Online Softmax,中间结果全部保留在 SRAM 里,只读写最终需要的 $Q$、$K$、$V$、$O$。
同一块 GPU、同一个模型、同一个数学定义,只改 kernel,就快了 2-4 倍、省了几十倍显存
这是过去五年 Transformer 加速里最有影响力的算法工程突破,现在是 PyTorch、vLLM、TensorRT-LLM 的默认实现。


10. 延伸阅读

  • 原论文:FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness(NeurIPS 2022)
  • 作者:Tri Dao 等(现 Together AI / Princeton)
  • GitHub:Dao-AILab/flash-attention