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
整个过程可以用三步概括:
- 打分:$S = QK^\mathsf{T}/\sqrt{d}$,得到 $N\times N$ 的分数矩阵
- 归一化:$P = \operatorname{softmax}(S)$,每行元素之和为 $1$
- 汇总:$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 之间来回搬数据:
- 计算 $S=QK^\mathsf{T}/\sqrt{d}$,再将 $N\times N$ 的 $S$ 写入 HBM。
- 从 HBM 读取 $S$,计算 $P=\operatorname{softmax}(S)$,再将 $P$ 写回 HBM。
- 从 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)}}.
$$
看着复杂,本质就三件事:
- 把旧的 $O$ 重新缩放,使用与更新 $\ell$ 相同的 $\exp\left(m^{(i-1)}-m^{(i)}\right)$ 修正因子
- 加上当前 tile 的贡献 $\exp\left(S^{(i)}-m^{(i)}\right)V^{(i)}$
- 使用新的 $\ell$ 归一化
6. FlashAttention 完整算法
6.1 核心思路
- Q 按行切块(外层循环)
- K, V 按行切块(内层循环)
- 每个 Q 块在 SRAM 里常驻,扫一遍所有 K/V 块
- 维护 m, ℓ, O 三个累积量,online 更新
- S 和 P 矩阵永远不完整出现,只有小 tile 大小
6.2 伪代码
1 | ## 输入:Q, K, V ∈ [N, d],都在 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+ | 4× 长 |
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