Published on

稀疏注意力

Authors
  • avatar
    Name
    Vegetog
    Twitter

稀疏注意力

一句话理解

稀疏注意力(Sparse Attention)的核心是:

不再让每个 Token 查看所有历史 Token,而只允许它查看一小部分“可能有用”的 Token。

它通过减少注意力连接数量,把长文本的计算量从平方级降到近似线性或次平方级。


1. 普通注意力为什么贵?

假设一句话有 nn 个 Token,每个 Token 会产生:

  • Query:我正在寻找什么;
  • Key:我这里有什么信息;
  • Value:如果你关注我,我能提供什么内容。

设:

Q,K,VRn×dQ,K,V\in\mathbb{R}^{n\times d}

标准注意力计算为:

Attention(Q,K,V)=softmax(QKTd)V\operatorname{Attention}(Q,K,V)=\operatorname{softmax}\left(\frac{QK^T}{\sqrt{d}}\right)V

其中最关键的是:

QKTRn×nQK^T\in\mathbb{R}^{n\times n}

它计算的是每个 Query 和每个 Key 的相关性。

假设有 4 个 Token:

我 喜欢 深度 学习

完整注意力关系是:

       我  喜欢  深度  学习
我     ●   ●    ●    ●
喜欢   ●   ●    ●    ●
深度   ●   ●    ●    ●
学习   ●   ●    ●    ●

一共需要计算:

4×4=164\times 4=16

如果有 128K Token,则理论上的注意力关系数量是:

1280002164 亿128000^2\approx 164\ \text{亿}

实际模型还有几十个注意力头、几十层,因此代价非常大。


2. 稀疏注意力究竟“稀疏”在哪里?

稀疏的不是 Q, K, VQ,\ K,\ V 本身,而是:

Query 和 Key 之间允许存在的连接。

可以引入一个掩码矩阵 MM

Attention(Q,K,V)=softmax(QKTd+M)V\operatorname{Attention}(Q,K,V)=\operatorname{softmax}\left(\frac{QK^T}{\sqrt{d}}+M\right)V

其中:

Mij={0,允许 Token i 查看 Token j,禁止查看M_{ij}=\begin{cases}0, & \text{允许 Token }i\text{ 查看 Token }j\\-\infty, & \text{禁止查看}\end{cases}

经过 Softmax 后,被设为 -\infty 的位置权重就变成 0。

但这里有一个非常关键的工程细节:

如果先完整计算 QKTQK^T,然后再把大部分位置遮掉,只减少了有效连接,没有减少主要计算量。

真正高效的稀疏注意力必须直接跳过无效区域,通常以 Token Block 为单位调用专门的稀疏 Kernel。


3. 最常见的滑动窗口注意力

滑动窗口注意力规定:每个 Token 只看附近的 ww 个 Token。

假设是因果语言模型,窗口大小 w=3w=3,包含当前 Token,那么:

Token 1 可以看:1
Token 2 可以看:1 2
Token 3 可以看:1 2 3
Token 4 可以看:2 3 4
Token 5 可以看:3 4 5
Token 6 可以看:4 5 6

对应注意力矩阵:

KeyQuery1 2 3 4 5 6

12        ● ●
3        ● ● ●
4          ● ● ●
5            ● ● ●
6              ● ● ●

原来每一行大约有 nn 个连接,现在每一行只有 ww 个连接:

O(n2)O(nw)O(n^2)\rightarrow O(nw)

如果 ww 是固定常数,那么相对于序列长度 nn,复杂度就是近似线性的:

O(n)O(n)

Longformer 就使用了“局部滑动窗口 + 少量全局注意力”的设计。Longformer 论文

窗口外的信息还能传过来吗?

可以间接传递。

假设每一层最多向前看 2 个 Token:

1层:Token 10 可以直接获得 Token 810 的信息
2层:Token 10 可以间接获得 Token 610 的信息
3层:Token 10 可以间接获得 Token 410 的信息

经过 LL 层,理论感受野大约可以扩大到:

L×(w1)L\times(w-1)

但这不等于远处信息被完整保留,因为它必须经过多个中间 Token 不断压缩和转发。

就像传话游戏:

虽然消息最终能传到很远,但经过很多人以后,细节可能已经丢失。


4. 常见的稀疏连接模式

4.1 局部窗口

每个 Token 看附近 Token

适合:

  • 语法和局部语义;
  • 代码中的相邻语句;
  • 图片中的附近像素;
  • 音频中的相邻时间片。

问题是跨文档的远距离关系传播较慢。


4.2 膨胀窗口

普通窗口连续查看附近位置:

i-3, i-2, i-1, i

膨胀窗口可以隔一定距离采样:

i-8, i-4, i-2, i

这样连接数量不变,但可以更快扩大感受野。

它类似于 CNN 里的空洞卷积:

不增加观察点数量,但把观察范围拉得更远。


4.3 全局 Token

选择少量特殊 Token,让它们和大量普通 Token 连接:

普通 Token:主要看附近
全局 Token:可以看整个文档

例如文档分类任务里,可以让 [CLS] 成为全局 Token:

普通 Token → 把局部信息汇总给 [CLS]
[CLS]      → 收集全文信息

在问答任务里,也可以让“问题部分”的 Token 成为全局 Token。

需要注意:

  • 在 BERT 这类双向 Encoder 中,全局 Token 可以双向查看全文。
  • 在自回归 Decoder 中仍然必须遵守因果性,前面的 Token 不能偷看未来。

4.4 随机连接

除局部连接外,再让每个 Token 随机连接少量远处 Token:

Token 100查看 96100
另外随机查看 174378

随机连接的作用类似于社交网络中的“远距离好友”:

大部分人只认识附近的人,但少量跨地区关系可以显著缩短信息传播路径。

BigBird 将三种模式结合起来:

  • 局部滑动窗口;
  • 少量随机连接;
  • 少量全局 Token。

在连接数近似线性的情况下,它仍然能形成较强的全局信息传播能力。BigBird 论文


4.5 分块稀疏注意力

GPU 并不擅长处理大量零散的单点连接,因此工程实现通常不会逐 Token 稀疏,而是按块处理。

例如每 128 个 Token 为一个 Block:

Block 1:Token 1128
Block 2:Token 129256
Block 3:Token 257384

然后规定:

Block 3 可以看 Block 2、Block 3
Block 4 可以看 Block 1、Block 3、Block 4

注意力矩阵就变成:

       K1 K2 K3 K4
Q1Q2     ●  ●
Q3        ●  ●
Q4     ●     ●  ●

这种方式叫 Block-Sparse Attention。

它可能会在一个有效 Block 内计算一些不重要的 Token 对,但换来的好处是:

  • 内存访问更规整;
  • 更容易使用矩阵乘法;
  • 更容易发挥 GPU Tensor Core;
  • 比不规则的逐元素稀疏更容易获得真实加速。

早期 Sparse Transformer 使用分解式稀疏模式,把复杂度从 O(n2)O(n^2) 降至 O(nn)O(n\sqrt{n})Sparse Transformer 论文


4.6 动态稀疏注意力

前面的窗口、全局和随机连接都是预先规定的。

动态稀疏注意力则根据内容选择:

当前 Query
寻找最相关的 Top-KKey
只对这些 Key 做完整注意力

听起来很理想,但有一个“先有鸡还是先有蛋”的问题:

如果要先计算所有 Query-Key 相似度才能选出 Top-K,那么最贵的 QKTQK^T 已经算完了。

因此,要真正节省计算,必须使用更便宜的候选选择方法,例如:

  • 聚类;
  • 哈希;
  • 路由网络;
  • 粗粒度 Block 打分;
  • 近似最近邻搜索。

动态稀疏更灵活,但会带来路由开销、不规则访存和负载不均衡。


5. Prefill 和 Decode 阶段有什么区别?

这是推理系统里非常重要的区别。

Prefill 阶段

假设 Prompt 有 nn 个 Token,需要同时计算它们之间的注意力。

完整注意力大约是:

O(n2)O(n^2)

如果每个 Token 只看 ww 个位置:

O(nw)O(nw)

所以稀疏注意力对长 Prompt 的 Prefill 特别有价值。

Decode 阶段

每次只生成一个新 Token,因此本轮 Query 长度是 1:

Query × 历史所有 Key

完整注意力单步复杂度约为:

O(n)O(n)

滑动窗口注意力单步只读取最近 ww 个 KV:

O(w)O(w)

如果模型所有相关层都保证只使用最近 ww 个 Token,那么超过窗口的旧 KV Cache 可以丢弃:

完整注意力 KV Cache:
K1 K2 K3 ... Kn 全部保存

滑动窗口 KV Cache:
只保存 K(n-w+1) ... Kn

因此稀疏注意力可能同时降低:

  • Decode 注意力计算量;
  • KV Cache 容量;
  • KV Cache 读取带宽。

但只有“逻辑上永远不会再次访问的 KV”才能安全删除。具有全局 Token、跨层混合或特殊检索路径的模型,需要保留相应的历史 KV。


6. 为什么理论上更省,实际却不一定更快?

假设:

完整注意力:计算100万个规则、连续的元素
稀疏注意力:计算20万个到处散落的元素

GPU 可能更喜欢前者,因为完整矩阵乘法:

  • 数据连续;
  • Tensor Core 利用率高;
  • Kernel 成熟;
  • 并行度高。

而不规则稀疏计算可能出现:

  • 需要读取稀疏索引;
  • GPU Warp 内线程工作量不同;
  • 内存访问不连续;
  • Kernel 启动和路由开销高;
  • 稀疏率不够高时,节省的 FLOPs 抵不过额外开销。

因此:

FLOPs 减少,不代表端到端延迟一定下降。

这也是为什么工程上更偏爱规则的 Block-Sparse,而不是随意的逐 Token 稀疏。


7. 它和 FlashAttention、PagedAttention 有什么区别?

这三个名字很容易混淆。

稀疏注意力

改变的是:

哪些 Query-Key 对需要计算。

它减少逻辑上的注意力连接和 FLOPs。

FlashAttention

不改变注意力结果,也不丢连接。

它依然计算完整的精确注意力,只是通过分块计算,减少 HBM 和片上 SRAM 之间的数据搬运,而且不需要在 HBM 中保存完整的 n×nn\times n 注意力矩阵。FlashAttention 论文

所以:

稀疏注意力:少算一些
FlashAttention:同样都算,但换一种更高效的计算顺序

二者也可以结合成 Block-Sparse FlashAttention。

PagedAttention

PagedAttention 主要解决:

KV Cache 怎么分配、映射、共享,减少显存碎片。

它把逻辑连续的 KV Cache 映射到可以不连续的物理 Block。

但它通常不会自动改变模型需要关注哪些历史 Token。对于完整注意力模型,历史 KV 逻辑上仍然需要保留;滑动窗口或稀疏模型才可能从模型结构上减少需要保存的历史 KV。


8. 稀疏注意力的主要缺点

远距离依赖可能被漏掉

例如第 10 个 Token 和第 100000 个 Token 强相关,但它们之间没有连接,模型就只能依赖多层间接传递。

稀疏规则需要和任务匹配

  • 文档适合局部窗口加全局 Token;
  • 图片可能适合二维窗口或轴向注意力;
  • 代码可能需要跨函数、跨文件连接;
  • 随机连接不一定符合真实语义。

没有一种稀疏模式对所有任务都最好。

不能随便在推理时替换

一个使用完整注意力训练的模型,如果推理时突然把大部分注意力连接删掉,分布发生变化,效果可能明显下降。

通常需要:

  • 从预训练阶段就使用该稀疏结构;或者
  • 对已有模型继续训练、微调,使它适应新的连接模式。

很长不等于真正理解

即使一个模型支持百万 Token:

  • 有效感受野可能没有覆盖全文;
  • 重要细节可能在多层传播中丢失;
  • 长上下文位置上的准确率可能不均匀;
  • “能放进去”不代表“能可靠找出来”。

最后的整体理解

稀疏注意力可以看作给 Token 之间修路:

完整注意力:
所有城市之间都有直达高速公路
道路太多,建设成本 O()

局部注意力:
只修相邻城市之间的道路
便宜,但跨国旅行要多次中转

全局 Token:
增加几个大型交通枢纽

随机注意力:
增加少量随机跨区域航线

Block-Sparse:
不是给每栋房子单独修路,而是在城市群之间修路
更适合 GPU 批量处理

最核心的三个结论是:

  1. 稀疏注意力减少的是 Token 之间的注意力连接。
  2. 真正加速必须在计算 QKTQK^T 时就跳过无效区域。
  3. 稀疏模式是在效率和信息可达性之间做权衡:连接越少越省,但越容易遗漏远距离信息。