← 返回题目列表

自注意力的计算复杂度是多少?FlashAttention 为什么能加速?

高频 困难 第 14 / 25 题 更新于 2026/08/03
注意力复杂度FlashAttention长序列IO-aware

简化版

自注意力对序列长度 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 三招:

  1. 分块(tiling):把 Q、K、V 切成小块,逐块加载进快得多的 SRAM(片上缓存)计算。
  2. 在线 softmax:softmax 需要全局最大值和归一化和,FlashAttention 用「边遍历边维护running最大值和running和」的数值稳定算法,分块累加也能得到正确结果,无需先算完整行
  3. 不落地中间矩阵:完整的 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 是固定的头维,真正随序列爆炸的是 。当上下文从 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 最大值 mrunning 归一化和 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)、速度提升数倍,且结果精确不近似。关键辨析两点:一是它慢的根因是访存受限而非算力;二是它降的是显存和墙钟时间,不降计算阶数(降阶要靠稀疏/线性注意力那类近似)。