高效注意力:线性与稀疏注意力变体

标准注意力是 Transformer 的能力来源,也是它的算力天花板。当序列长度增加时,注意力的时间和显存开销会按序列长度的平方增长,这直接限制了模型能处理的上下文长度。本文先拆解 O(n²) 瓶颈的来由,再讲两类主流解法——稀疏注意力与线性注意力——各自如何把复杂度降到近线性,最后回顾 FlashAttention 的「省显存」思路,并给出长序列场景下的选型建议。本文侧重直觉与思路,不要求严格数学推导;需要精确公式与实验细节,可查阅文末参考。

标准注意力的 O(n²) 瓶颈

设输入序列长度为 n,每个 token 的表示维度为 d。我们先用输入 X 投影出查询 Q、键 K、值 V(均为 n×d 的矩阵),标准缩放点积注意力的计算如下:

注意力分数 S = Q Kᵀ       形状为 n × n
注意力权重 A = softmax(S)  形状仍为 n × n
输出      O = A V         形状为 n × d

瓶颈出现在两处:

第一是计算。算出 Q Kᵀ 需要 n×n×d 次乘加,对 n 个查询各自做 softmax 也是 O(n²) 量级,最后与 V 相乘又是 n×n×d。只要 n 变大,开销就随 n 的平方膨胀。

第二是显存。标准实现会把完整的 n×n 注意力权重矩阵 A 物化(materialize)到显存里。当 n 等于 4096 时,n² 约为 1600 万;当 n 等于 65536 时,n² 约为 43 亿。再乘上批大小与注意力头数,显存很快被吃满。这正是「序列一长就爆显存」的根因。

下面用一张表直观对比不同序列长度下的注意力矩阵规模(以单头、不乘批大小计):

序列长度 n    注意力矩阵 n×n    约需元素数
   512          512 × 512       约 26 万
   4096        4096 × 4096      约 1600 万
   16384      16384 × 16384     约 2.7 亿
   65536      65536 × 65536     约 43 亿

注意,这段表里的具体数字仅用于说明平方增长的趋势,真实显存占用还取决于精度与框架实现细节(待核实)。

稀疏注意力:只算该算的那些

稀疏注意力的核心直觉很简单:并非每个 token 都需要和全部 token 交互。通过人为规定「每个位置只关注一小部分位置」,就能避免生成完整的 n×n 矩阵,把开销从平方压到近线性。下面看两个代表性工作。

Longformer:滑动窗口加全局注意力

Longformer(Beltagy、Peters、Cohan,arXiv:2004.05150,已核验)提出一种随序列长度线性增长的注意力机制,可作为标准自注意力的直接替换。它把两种注意力组合在一起:

  • 局部窗口注意力:每个 token 只关注左右各 w 个邻居,相当于一个滑动窗口。每个位置只看固定数量的上下文,整体复杂度约为 O(n·w),即随 n 线性增长。
  • 全局注意力:针对任务挑选少数几个「全局 token」(例如分类用的 CLS、问答中的问句 token),让它们能关注整条序列,也被整条序列关注,用来承载必须全局传递的信息。

论文还使用了空洞滑动窗口(dilated sliding window),让窗口之间留空以扩大感受野而不增加计算量。据论文摘要,Longformer 在字符级语言建模(text8、enwik8)上取得当时最佳结果,在长文档任务上稳定超过 RoBERTa,并在 WikiHop、TriviaQA 上刷新当时最佳(待核实)。论文后续版本还给出了面向长文档生成任务的长文档编码器-解码器变体 LED。

BigBird:随机、窗口、全局三合一

BigBird(Zaheer 等人,NeurIPS 2020,arXiv:2007.14062,已核验)把稀疏模式拆成三类,再拼到一起:

  • 随机注意力:每个 token 随机挑选少量其他 token 建立连接,用来传播非局部信息。
  • 窗口注意力:和 Longformer 类似,关注局部邻居。
  • 全局注意力:少数全局 token 负责跨序列的全局汇聚。

三类注意力加起来的总连接数约为 O(n·(r + w + g)),其中 r、w、g 分别是随机、窗口、全局的连接数,均为与 n 无关的常数,因此整体复杂度是线性的。论文的理论分析指出,BigBird 是序列函数的通用逼近器且是图灵完备的,因此保留了完整注意力模型的关键性质。据摘要,借助稀疏注意力,BigBird 能处理的序列长度可达此前同类硬件上限的约 8 倍(待核实),并在问答、摘要等任务上明显提升,还被尝试用于基因组数据。

常见的稀疏模式

把上面的设计抽象一下,稀疏注意力大致有以下几种「图案」:

滑动窗口(window):只看邻近位置,感受野有限
空洞窗口(dilated):窗口间留空,扩大感受野
分块局部(block):把序列切成块,块内全连接
跨步(strided):按固定步长取位置
随机(random):随机连少量边,传播长程信号
全局(global):少数 token 与全体相连

稀疏注意力的代价是「近似」:被丢掉的连线意味着信息损失,效果依赖稀疏图案是否与任务匹配,窗口大小、全局 token 的选择也都需要调参,实现上也比标准注意力复杂。

线性注意力:用核函数改写公式

稀疏注意力是从「减少连接数」入手,线性注意力则换了一条路:直接改写注意力的数学形式,让那个 n×n 的矩阵根本不必出现。

标准注意力可以写成:

O = softmax(Q Kᵀ) V

线性注意力的思路是,把 softmax 换成某个核函数对应的特征映射 φ,将上式改写为:

O = φ(Q) (φ(K)ᵀ V)

关键在括号的位置变了。标准写法先算 Q Kᵀ 得到 n×n,再乘 V;改写后先算 φ(K)ᵀ V。我们看一下形状:φ(K) 是 n×d’ 的矩阵,V 是 n×d 的矩阵,φ(K)ᵀ V 的结果是 d’×d——它的大小只和特征维度 d’、d 有关,与序列长度 n 完全无关。之后再让 φ(Q)(n×d’)去乘这个 d’×d 的中间结果,得到 n×d 的输出。整个过程再也没有形成 n×n 的矩阵,复杂度从 O(n²) 降到了 O(n·d·d’),即随 n 线性增长。

这一步之所以成立,依赖矩阵乘法的结合律:

( φ(Q) φ(K)ᵀ ) V  =  φ(Q) ( φ(K)ᵀ V )

左边是先凑出 n×n 再乘 V,右边是先凑出与 n 无关的 d’×d 再乘 φ(Q)。线性注意力就是利用了「先算小的那一项」来避开平方开销。这一思路由 Katharopoulos 等人(ICML 2020,arXiv:2006.16236,已核验)系统提出,他们把自注意力表达为核特征映射的线性点积,并指出这种形式可以迭代实现,从而揭示注意力与循环神经网络之间的关系;据摘要,线性注意力在自回归生成长序列时可比原始 Transformer 快至多约 4000 倍(待核实)。

下面用一段精简的 PyTorch 风格伪代码展示「先算小的那一项」:

import torch

def linear_attention(q, k, v, feature_map):
    # q, k, v 形状均为 (batch, n, d)
    # feature_map 即论文中的 φ,例如 elu(x) + 1
    phi_q = feature_map(q)  # (batch, n, d')
    phi_k = feature_map(k)  # (batch, n, d')

    # 先算与 n 无关的中间量:对 n 求和
    kv = torch.einsum('bnd,bne->bde', phi_k, v)  # (batch, d', d)
    # 再乘回查询,得到逐位置输出
    out = torch.einsum('bnd,bde->bne', phi_q, kv)  # (batch, n, d)

    # 实际实现还需除以归一化项 sum(phi_k),此处省略以突出核心思路
    return out

线性注意力的代价也在「近似」:特征映射 φ 无法完全等价于 softmax,表达力通常弱于标准注意力,效果好坏与 φ 的选择强相关;在某些任务上它可能不如 softmax 注意力精确。

FlashAttention 回顾:为什么省显存

前序教程「注意力机制深入」已经讲过 FlashAttention 的基本机制,这里只补一个视角——它为什么能省显存。要点是:FlashAttention 不是近似注意力,而是精确的注意力,它通过「IO 感知(IO-aware)」来同时提速与省显存。

标准注意力之所以吃显存,是因为它把完整的 n×n 注意力权重矩阵写到了 GPU 的高带宽显存(HBM)里,之后又要反复从 HBM 读回来做 softmax 和加权求和。HBM 容量大但速度慢,频繁的大量读写既占显存又拖慢速度。

FlashAttention(Dao 等人,arXiv:2205.14135,已核验)采用分块(tiling)策略:把 Q、K、V 切成小块,在速度极快但容量很小的片上 SRAM 中完成每一块的注意力计算,只把最终结果写回 HBM,而从不把完整的 n×n 矩阵物化出来。因为中间的大矩阵始终留在 SRAM、不落 HBM,显存占用从 O(n²) 降到 O(n),同时 HBM 的读写次数大幅减少,速度也随之提升。论文的 IO 复杂度分析表明,FlashAttention 比标准注意力需要更少的 HBM 访问,并在一定 SRAM 尺寸范围内是最优的。

这里要厘清一个常见误解:FlashAttention 的计算量在理论上仍是 O(n²)(它是精确注意力,没有省略任何计算),它真正压下来的是显存和 IO 开销,而不是把计算复杂度降到线性。据摘要,FlashAttention 在 BERT-large(序列长度 512)上取得约 15% 的端到端加速,在 GPT-2(序列长度 1K)上约 3 倍加速,在长程基准上约 2.4 倍加速,并首次让 Transformer 在序列长度 16K 的 Path-X、序列长度 64K 的 Path-256 上取得好于随机的表现(待核实)。

选型:长序列场景该用哪种

面对长序列,先分清两件事:你能否接受「近似」,以及你卡在的是「显存」还是「计算」。下面给出一张速查表。

场景                               推荐做法
短序列(约 512 以内)             标准注意力即可;想省显存提速可用 FlashAttention
长序列、要保精确结果             FlashAttention(及块稀疏 FlashAttention)
超长序列、可接受近似(编码类)   稀疏注意力,如 Longformer、BigBird
流式/自回归、要常显存解码         线性注意力,可像 RNN 一样以常数显存逐步生成

各自的代价汇总如下:

  • 标准注意力:实现简单、效果稳,但序列一长就爆显存、算得慢,不适合长上下文。
  • FlashAttention:结果精确、无质量损失,且显存降到线性、速度更快;但计算量仍是平方级,极端长序列下训练耗时依然可观,且依赖特定硬件与算子实现。
  • 稀疏注意力:复杂度降到近线性,适合长文档编码;代价是近似带来潜在质量损失,稀疏图案和全局 token 需要按任务调参,实现复杂。
  • 线性注意力:复杂度真正线性,自回归时显存为常数,适合流式与超长序列;代价是特征映射的近似可能削弱表达力,部分任务上不及 softmax 注意力。

一句话原则:要精确且省显存,优先 FlashAttention;要进一步把计算也压到线性且能容忍近似,再在稀疏与线性之间按任务形态(编码还是生成、是否流式)取舍。

小结

  • 标准注意力会物化 n×n 的注意力矩阵,时间与显存均随序列长度平方增长,这是长上下文的根本瓶颈。
  • 稀疏注意力(Longformer 的滑动窗口加全局、BigBird 的随机加窗口加全局)通过只计算部分连接,把开销降到近线性,代价是引入了近似。
  • 线性注意力用核函数特征映射 φ 改写 softmax(Q Kᵀ) V 为 φ(Q)(φ(K)ᵀ V),依靠矩阵乘法结合律先算与 n 无关的小项,使复杂度降到 O(n),代价同样是近似且依赖 φ 的选择。
  • FlashAttention 不改计算结果(精确注意力),而是用 IO 感知的分块策略避免把大矩阵写入慢速 HBM,从而把显存降到 O(n) 并提速;它的计算量仍是平方级。
  • 选型时先问「能否接受近似」和「卡在显存还是计算」:精确省显存用 FlashAttention,近线性计算且可近似再选稀疏或线性。

参考与延伸阅读

  • Longformer: The Long-Document Transformer(Beltagy、Peters、Cohan,2020)。arXiv:2004.05150。已核验:WebFetch 摘要核验,确认其采用局部窗口注意力与任务驱动的全局注意力、复杂度随序列长度线性增长。(具体基准指标如 WikiHop、TriviaQA 提升幅度待核实)
  • Big Bird: Transformers for Longer Sequences(Zaheer 等人,NeurIPS 2020)。arXiv:2007.14062。已核验:WebFetch 摘要核验,确认其由随机、窗口、全局三类稀疏注意力组合而成,复杂度降为线性,并号称可处理约 8 倍长的序列。(「8 倍」「图灵完备」等具体论断待核实)
  • Transformers are RNNs: Fast Autoregressive Transformers with Linear Attention(Katharopoulos 等人,ICML 2020)。arXiv:2006.16236。已核验:WebFetch 摘要核验,确认其将自注意力表达为核特征映射的线性点积,借助矩阵乘法结合律把复杂度从 O(N²) 降到 O(N)。(「快至多 4000 倍」等具体指标待核实)
  • FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness(Dao 等人,2022)。arXiv:2205.14135。已核验:WebFetch 摘要核验,确认其通过分块减少 HBM 与 SRAM 间读写、为 IO 感知的精确注意力,显存占用从平方降为线性。(BERT-large 15%、GPT-2 3 倍、Path-X 16K 等具体指标待核实)
  • 前序教程「注意力机制深入:从加性注意力到多头与 FlashAttention」(同站 attention-deep-dive)。已核验:站内存档,标准注意力、缩放点积、多头与 FlashAttention 基础背景可在此查阅。
本文累计阅读