FlashAttention: attention 为什么还能更快?
这一节单独把 FlashAttention 拿出来讲,因为它特别容易让人误会成“另一种 attention 算法”。
但它真正重要的地方,其实不是改公式,而是改实现。
这一节准备回答几个问题:
- FlashAttention 到底在解决什么问题?
- 它有没有改变 attention 的数学定义?
- 为什么 attention 明明公式很清楚,工程上还是会很慢、很吃显存?
- 理解 FlashAttention 时,最值得记住的结论是什么?
Q1: FlashAttention 到底在解决什么问题?
先回到最普通的 scaled dot-product attention:
设
- \( Q \in \mathbb{R}^{n \times d_k} \)
- \( K \in \mathbb{R}^{n \times d_k} \)
- \( V \in \mathbb{R}^{n \times d_v} \)
那么 attention 写成:
$$ \mathrm{Attention}(Q,K,V)=\mathrm{softmax}\left(\frac{QK^\top}{\sqrt{d_k}}\right)V $$
这套数学定义本身没有问题,问题主要出在实现上。
因为如果直接按这个式子展开,往往会显式构造一个 shape 为 \( n \times n \) 的 attention score 矩阵:
$$ S = QK^\top \in \mathbb{R}^{n \times n} $$
当序列一长,这个中间结果就会很大:
- 显存占用会迅速上升
- 读写这个大矩阵本身也会变慢
- 训练和推理都容易被 memory bandwidth 卡住
所以 FlashAttention 主要解决的不是“attention 会不会算错”,而是:
- attention 能不能少存一些中间结果
- attention 能不能少做一些低效的显存读写
- 在不改数学结果的前提下,把实现做得更贴近 GPU
Q2: 它有没有改变 attention 的数学定义?
通常没有。
这是理解 FlashAttention 时最重要的一点之一。
它并不是把 attention 换成了另一个近似很强的新公式,而是尽量保持输出和标准 attention 一致,只是换了一种更高效的计算路径。
也就是说,下面这个目标没变:
$$ \mathrm{softmax}\left(\frac{QK^\top}{\sqrt{d_k}}\right)V $$
变的是“如何得到它”:
- 不再粗暴地把所有中间矩阵一次性完整落到显存里
- 而是更倾向于分块计算、边算边归约、尽量减少中间结果回写
所以一个更准确的说法是:
- FlashAttention 首先是 attention 的高效实现方案
- 而不是一种彻底改写定义的注意力机制
Q3: 为什么 attention 明明公式很清楚,工程上还是会很慢、很吃显存?
因为真正贵的,常常不是公式看起来有多复杂,而是中间张量有多大、数据搬运有多频繁。
以最基本的 attention 为例,中间会涉及:
- \( QK^\top \in \mathbb{R}^{n \times n} \)
- softmax 后的权重矩阵 \( A \in \mathbb{R}^{n \times n} \)
- 再和 \( V \) 相乘得到输出
所以当 \( n \) 变大时,问题不只是计算量接近 \( O(n^2 d) \),还包括:
- 中间激活会很大
- 显存读写会变重
- kernel 之间来回搬运数据会带来额外开销
这也是为什么很多 attention 优化工作,看起来都像是在“做实现细节”,但实际上影响很大。
因为在大模型里,底层实现细节本身就会直接决定吞吐、显存和可训练长度。
Q4: 理解 FlashAttention 时,最值得记住的结论是什么?
如果只保留最重要的几点,我觉得是:
- FlashAttention 主要优化的是 attention 的实现方式,而不是它的基本数学定义。
- 它之所以重要,是因为标准 attention 的中间结果很大,显存和带宽会很快成为瓶颈。
- 理解它最好的角度不是“新机制”,而是“更高效地把同一个机制跑起来”。
所以如果后面在代码里看到 flash_attn 这一类实现,第一反应不应该是“模型换结构了”,而更应该是:
- 这通常是在解决 attention 太慢、太吃显存的问题