跳转至

效率型注意力

稀疏注意力主要解决什么问题?它基于什么假设?

解决的问题: 标准自注意力机制的计算复杂度为 O(N2d),其中 NN 是序列长度,d 是每个头的维度。当 N 达到数千或数万时(如长文档、高分辨率图像、DNA序列),计算每个 token 对之间的注意力分数将产生巨大的计算开销和显存占用。具体来说,存储N×N 的注意力矩阵需要O(N2) 内存,反向传播时需要缓存该矩阵或重计算,显存极易溢出。这严重限制了 Transformer 处理长序列的能力。

稀疏注意力旨在将这一复杂度降低到 O(Nkd) 或 O(Nd2),其中 kN 是每个 token 关注的固定数量的其他 token。通过限制每个查询 token 只能与一个稀疏的 key 集合交互,避免了构建稠密的 N2 矩阵。

核心假设:

稀疏注意力建立在这样一个先验假设之上:注意力矩阵是高度稀疏的,且大部分有价值的信息交互集中在局部或特定的全局节点上。

具体而言:

  • 局部性假设:自然语言、图像、音频等数据具有很强的局部相关性。一个词的语义通常由其邻近的词决定,远处的直接关联很弱(除非有特殊指代)。

  • 全局节点的桥梁作用:某些特殊 token(如 [CLS]、句号、段落起始标记)或依据任务选定的 token 承载了汇聚和分发全局信息的功能,应当被所有 token 关注。

  • 随机连接补充:类似随机图的性质,在局部和全局之外,少量随机连接能保证信息快速混合,使得任意两个 token 之间的信号传递路径长度维持在O(logN) 级别,并能从理论上逼近全注意力。

基于此,可以设计出多种稀疏模式(局部窗口、空洞窗口、全局注意力、随机注意力等),在极大降低计算量的同时,尽可能保持模型的长期依赖建模能力。


Longformer 的注意力模式是怎样的?滑窗+全局注意力如何组合?

Longformer 的注意力模式是由 滑动窗口注意力(Sliding Window Attention) 和 全局注意力(Global Attention) 两个部分组合而成的混合稀疏注意力。

image.png

全局注意力

  • 预先选择少量的 token 作为“全局 token”。例如,分类任务中的 [CLS] token,问答任务中问题段落的分隔符,或者所有标点符号等。

  • 全局 token 分两类操作:

  • 作为查询(Query):全局 token 可以关注序列中的所有 token,从而汇聚整个输入的信息。
  • 作为键(Key):所有普通 token 也可以关注这些全局 token,使得全局信息能够广播到序列的每个位置。

  • 通常全局 token 的数量 gg 非常小(例如 1~100),所以这部分额外开销为 O(Ng),同样线性。

组合方式

在具体计算中,Longformer 为注意力矩阵设计了一个掩码模式:

  • 对于每个非全局 token,掩码仅允许其关注窗口内的 token 和所有全局 token。

  • 对于每个全局 token,掩码允许其关注所有 token(即稠密行)。

这样,注意力矩阵被分解为一个稠密条带(局部)加上少数全连接的行和列(全局)。这种组合使得每个 token 既能通过局部窗口维持细节上下文,又能通过全局 token 在长距离上迅速传递信息,无需为整个 N×N 矩阵付出代价。


BigBird 是如何实现随机、局部和全局注意力混合的?复杂度可降到多少?

BigBird 证明了一个重要定理:通过结合三种稀疏注意力模式,模型既能保持全注意力的表达能力(如图灵完备性),又能将复杂度降至线性。其混合模式包括:

局部窗口注意力

image.png

全局 token

选择一些 token 作为全局节点(例如 [CLS] 以及每 d 个 token 选一个,或随机选取),总数 g 为常数。这些 token 参与全注意力的行和列。

随机注意力

每个 token 额外随机选择 r 个其他 token 参与注意力计算。这些随机连接可以看作在局部图和全局节点之间增加了随机边,使图变成扩展图(expander graph),信息混合极快。 BigBird 理论证明,在局部窗口、全局节点和随机边的共同作用下,注意力图成为一个扩张性极好的随机图,任意两个节点之间的路径长度很短,使得 Transformer 在理论上仍然是通用的近似器。

复杂度

每个 token 关注的键数量为:w(局部)+ g(全局)+ r(随机)。所有这些参数都是常数(与序列长度 N 无关),因此计算和内存复杂度为 O(N),即线性复杂度。

实际上,BigBird 在长序列(如 4096 或更长)上的训练和推理速度比全注意力快数倍,且显存占用大幅降低,同时在长文档摘要、问答等任务上保持了与全注意力相当的性能。


线性注意力如何将 softmax 注意力的指数形式转化为核函数形式?以 Performer 为例说明。

标准注意力:

image.png

ϕ(q)⊤ϕ(k) 是 exp(qk) 的无偏估计。实际中还常对 W 进行正交化以降低方差,即 FAVOR+ 中的 O(正交化)和 +(保证正值),进一步提高近似质量和训练稳定性。这种映射将指数形式的 softmax 注意力转化成了可分解为内积的线性核形式,从而可以使用结合律将复杂度降为线性。


Performer 的 FAVOR+ 算法是如何实现线性注意力的?为什么能保证无偏估计?

FAVOR+ 全称是 Fast Attention Via Orthogonal Random features (with "+" for positive features)。它正是上一节描述的通过随机特征映射实现线性注意力的具体算法。其实现步骤和原理如下:

实现步骤

image.png

为什么是无偏估计?

从数学上看,对于任意查询 q 和键 k,因为 ω∼N(0,I),有:

image.png

image.png


Linformer 如何通过低秩分解将注意力复杂度降为线性?它的投影矩阵有什么限制?

低秩分解降复杂度

Linformer 的关键发现是:经过 softmax 归一化后的注意力矩阵 P(形状 N×N)在真实数据上通常是低秩的。因此,可以直接在 K 和 V 的序列长度维度上应用线性投影,将其压缩到一个固定的较小维度 k(如 128 或 256)。

image.png

投影矩阵的限制

  1. 依赖序列长度 N:投影矩阵 E,F 的大小是 k×N,其中 N 是训练时的最大序列长度。这导致模型无法直接应用于比训练长度更长的序列,除非对投影矩阵进行插值、截断或重新设计。

  2. 泛化到变长序列的困难:实际中可以通过参数化的投影(如用卷积层或多层感知机)来替代固定的矩阵,使其能处理可变长度。例如,在 Linformer 的实践中,使用了层级的卷积投影,这样在推理时可以动态适应序列长度,但性能可能与固定 N 的矩阵有差异。

  3. 线性投影可能损失信息:将 N 维压缩到 k 维,如果 k 太小,会丢失细粒度的 token 差异,影响需要精确 token 级别的任务(如问答定位)。但对于很多全局理解任务(如分类、摘要)影响较小。

  4. 跨序列共享:投影在训练和推理时是共享的,但若测试时的序列长度与训练差异极大,投影的有效性会下降。


线性注意力的主要缺点是什么?在哪些任务上可能效果不如标准注意力?

主要缺点

  • 失去 softmax 的“稀疏选择性”:标准 softmax 将注意力权重挤压到局部高点,极大地抑制无关 token,使得模型能聚焦于最关键的信息。线性注意力本质上是核平滑,注意力权重更“平均”,缺少这种尖锐的聚焦能力,容易让大量无关信息混入聚合。

image.png

效果可能不如标准注意力的任务

  • 精确的长文档问答:需要在长文中定位某一个小片段的证据,标准注意力或窗口+全局的稀疏注意力能直接关注到相关 token,而线性注意力的聚合表示可能无法保留如此精细的 token 级区别。

  • 需要强对齐的任务:如跨句子语义匹配、共指消解,这些任务需要计算特定 token 对之间的精确相似度,线性核近似可能不够精确。

  • 少样本学习(Few-shot):在上下文学习(in-context learning)中,标准注意力的稀疏性质有助于从示例中迅速抓取模式,而线性注意力可能将示例的细节过度平滑,降低学习效率。

  • 生成任务中的复制能力:线性注意力在记忆和复制长序列中的特定 token 时往往不如稠密注意力或局部+全局注意力可靠。


FlashAttention 实现了“IO感知”,具体是如何通过分块和重计算来加速的?

IO 感知(IO-aware) 指的是算法设计时充分考虑了现代 GPU 的内存层次结构:

  • 高带宽显存(HBM):容量大(如 80GB),但带宽相对较低(约 2 TB/s),且延迟高。

  • 片上 SRAM(Static Random-Access Memory):每个流多处理器(SM)上有一小块(如 192 KB),带宽极高(约 19 TB/s)但容量极小。

标准注意力中,计算 S=QK⊤ 需要将巨大的中间矩阵 S(N×N)写入 HBM 再读出以进行 softmax 和乘以 V。这种对 HBM 的频繁读写成为瓶颈(计算是快速的,数据搬运是慢的)。

FlashAttention 的加速策略:

image.png

通过这些技术,FlashAttention 实现了在长序列上 2~4 倍的训练加速,并将显存占用降低 5~10 倍,同时保持数学上的精确等价(无近似)。


FlashAttention 的前向和反向传播是如何在 SRAM 中分块计算的?为什么能节省显存?

前向传播的分块计算

image.png

反向传播的分块计算

反向传播需要梯度 dQ,dK,dV。由于没有保存注意力矩阵 P,需要重计算:

image.png

节省显存的原因

标准实现需要在前向存储完整的 N×N 注意力矩阵 P 用于反向,这对长序列是巨大的显存开销(例如 N=8k, 每个头 FP16 下约 128 MB,多头则更大)。FlashAttention 仅需存储 O(N) 的统计量(m, l)和输出 O,反向时重新计算 P。尽管重计算带来约 30% 的额外计算量,但由于 GPU 瓶颈在内存带宽而非计算,总的运行时间反而下降,且显存占用降至原来的 1/N 量级,使得 batch size 可增大,进一步提升了训练效率和模型扩展能力。

这九个问题的详细解答覆盖了从理论假设到具体算法实现的各个层面,希望能满足你对深度的要求。