自注意力的计算复杂度是多少?FlashAttention 为什么能加速?
简化版
自注意力对序列长度 n 的计算和显存都是 O(n²)——要算 n×n 的注意力分数矩阵,这是长序列的主要瓶颈。FlashAttention 是一种 IO 感知的精确注意力加速法:它不把巨大的 n×n 矩阵写进显存(HBM),而是把 Q、K、V 分块(tiling)加载到芯片上的高速缓存 SRAM 里,用在线 softmax(online softmax) 分块累积地算出结果,全程不落地完整注意力矩阵。这样把显存从 O(n²) 降到 O(n),并通过大幅减少 HBM 读写把速度提升数倍——而且结果是精确的,不是近似。
详细版
复杂度从哪来
标准注意力:
S = Q·Kᵀ # (n×d)·(d×n) = n×n, 计算 O(n²·d)
P = softmax(S) # n×n
O = P·V # (n×n)·(n×d), 计算 O(n²·d)
- 计算复杂度 O(n²·d):两次 n×n 规模的矩阵乘。
- 显存复杂度 O(n²):要存下 n×n 的 S 和 P。序列一长(如 n=32K),n² 达十亿级,显存和访存都爆炸。
FlashAttention 的关键洞察
标准实现慢的真正原因不是算力不够,而是访存(IO)太多:n×n 的中间矩阵要在慢速显存 HBM 和计算单元之间反复读写。GPU 的算力远快于 HBM 带宽,注意力是**访存受限(memory-bound)**的。
FlashAttention 三招:
- 分块(tiling):把 Q、K、V 切成小块,逐块加载进快得多的 SRAM(片上缓存)计算。
- 在线 softmax:softmax 需要全局最大值和归一化和,FlashAttention 用「边遍历边维护running最大值和running和」的数值稳定算法,分块累加也能得到正确结果,无需先算完整行。
- 不落地中间矩阵:完整的 n×n 分数矩阵从不写回 HBM,显存降到 O(n);反向传播时靠重计算(recompute)而非存储中间矩阵。
收益
- 显存:O(n²) → O(n),能训更长上下文。
- 速度:大幅减少 HBM 读写,实测 2–4 倍甚至更高。
- 精确:数学上等价于标准注意力,不是稀疏/线性近似。
完整版教学
一、把 O(n²) 说清楚:贵在哪一步
注意力要让每个 token 和其他所有 token 算相关度,于是天然产生一个 n×n 的分数矩阵 S=QKᵀ。
- 算 S:n×n 个分数,每个是 d 维点积,O(n²·d) 计算。
- softmax(S):对 n×n 逐行归一化。
- 乘 V:又是 O(n²·d)。
所以计算是 O(n²·d)、显存是 O(n²)(存 S 和 P)。d 是固定的头维,真正随序列爆炸的是 n²。当上下文从 2K 涨到 32K,n² 涨 256 倍——这就是长上下文又慢又占显存的根源,也是各种「高效注意力」要攻克的目标。
二、一个反直觉的事实:注意力慢,主要慢在搬数据
很多人以为注意力慢是「乘法太多、算力不够」。FlashAttention 的作者指出真相:在现代 GPU 上,注意力是访存受限(memory-bound)而非计算受限。
原因是 GPU 有两级存储:
- HBM(高带宽显存):容量大(几十 GB)但带宽相对慢。
- SRAM(片上缓存):极快但容量极小(每个 SM 只有几十 KB~几 MB)。
标准注意力实现会把 n×n 的 S、P 反复在 HBM 里读写:算 S 写回 HBM,读回来做 softmax 写回 HBM,再读回来乘 V……海量的 HBM 读写成了瓶颈,而 GPU 的计算单元大量时间在「等数据」。所以优化的关键不是减少乘法,而是减少 HBM 访问。
三、FlashAttention 三板斧详解
① 分块(Tiling)——把计算搬进 SRAM 把 Q、K、V 沿序列维切成小块。外层循环遍历 Q 的块,内层循环遍历 K/V 的块,每次只把当前小块加载进 SRAM 计算局部注意力。因为块很小,能装进快速的 SRAM,避免把大矩阵放 HBM。
② 在线 softmax(Online Softmax)——分块也能算对归一化 难点在于 softmax 需要「整行的最大值和求和」来归一化,而分块时看不到整行。FlashAttention 借用「在线/流式 softmax」技巧:遍历每个 K/V 块时,维护当前的running 最大值 m 和 running 归一化和 l,每来一个新块就用「重新缩放旧结果 + 累加新块」的方式更新,数学上保证最终等于对整行做 softmax。这样不必先看完整行,边走边算就得到精确结果。(减去最大值是为了数值稳定,防止 exp 溢出。)
③ 不落地 + 反向重计算——省显存 前向从不把完整 n×n 矩阵写回 HBM,显存需求降到 O(n)。反向传播本来需要中间的 P 矩阵,FlashAttention 选择不存它、需要时按块重新算(recomputation)——用一点额外计算换巨大的显存节省,而由于瓶颈本就在访存,这笔交换非常划算。
简要说抓住精髓:FlashAttention 不减少浮点运算,而是通过”分块进 SRAM + 在线 softmax + 不落地大矩阵”把慢速 HBM 的读写砍掉一个量级,从而又快又省显存,且结果精确。
四、FlashAttention vs 其他”高效注意力”
历史上为对付 O(n²) 有很多路线,FlashAttention 属于「精确、IO 优化」这一支,和「近似」路线本质不同:
| 路线 | 代表 | 是否精确 | 思路 |
|---|---|---|---|
| 稀疏注意力 | Longformer、BigBird | 近似 | 只算局部+少量全局,跳过大部分对 |
| 线性注意力 | Performer、Linear Transformer | 近似 | 用核技巧把 O(n²) 降到 O(n) |
| 低秩/压缩 | Linformer | 近似 | 把 K/V 投影到低维 |
| IO 感知精确 | FlashAttention | 精确 | 不改数学,优化访存与分块 |
近似方法能降复杂度阶,但常有质量损失或适用局限;FlashAttention 保持与标准注意力完全一致的结果,又拿到大幅加速和显存节省,因此被主流框架(PyTorch SDPA、vLLM 等)默认采用。FlashAttention-2、-3 进一步优化并行度和对新硬件的适配。
五、复杂度到底降没降
要点清楚:FlashAttention 的计算复杂度仍是 O(n²·d)(浮点运算数没变),它降的是显存复杂度(O(n²)→O(n))和实际墙钟时间(靠减少 HBM 访问)。想在计算阶数上降到 O(n) 得靠稀疏/线性注意力那类近似方法。这个区别是高频考点,别答错。
六、常见误区与追问
- 误区:FlashAttention 降低了计算复杂度。 没有,计算仍 O(n²);它降的是显存和访存,靠 IO 优化提速,结果精确。
- 误区:它是一种近似注意力。 它是精确的,和标准注意力数学等价。
- 追问:为什么标准注意力慢? 访存受限——n×n 中间矩阵在 HBM 反复读写,GPU 算力在等数据。
- 追问:在线 softmax 解决什么? 让分块计算也能得到正确的全局 softmax 归一化,无需先算完整行。
- 追问:反向为什么要重计算? 前向没存 n×n 中间矩阵以省显存,反向需要时按块重算,用算力换显存。
- 追问:它和 GQA/MQA 冲突吗? 不冲突,正交——FlashAttention 优化注意力的算法/访存,GQA/MQA 优化 KV 存储量,常一起用。
- 追问:还有哪些降长序列成本的手段? 稀疏/线性注意力(降阶但近似)、滑动窗口、MoE(与序列无关)、KV 压缩等。
七、加强记忆
复杂度记死:自注意力计算 O(n²·d)、显存 O(n²),长序列瓶颈在 n²。 FlashAttention 记「IO 感知、精确、三板斧」:分块进 SRAM + 在线 softmax + 不落地大矩阵(反向重计算),把慢速 HBM 的读写砍掉一个量级,显存降到 O(n)、速度提升数倍,且结果精确不近似。关键辨析两点:一是它慢的根因是访存受限而非算力;二是它降的是显存和墙钟时间,不降计算阶数(降阶要靠稀疏/线性注意力那类近似)。