重读 Attention:从 O(n²) 到线性的六种妥协

把 Linformer、Performer、Mamba、滑动窗口放在同一张坐标系里比较,看它们各自放弃了什么样的长程依赖能力。

每隔几个月就会有一篇”高效注意力”的论文声称把复杂度降到线性,同时”性能几乎无损”。读得多了会发现,这句话的重点全在”几乎”上。

标准注意力的计算量来自那个 n×nn \times n 的相似度矩阵:

Attn(Q,K,V)=softmax ⁣(QKdk)V\text{Attn}(Q, K, V) = \operatorname{softmax}\!\left(\frac{QK^\top}{\sqrt{d_k}}\right) V

想让它变快,只有两条路:要么别算全部的格子,要么别显式构造这个矩阵。所有方案都是这两条路的变体。

两条路径的分野

稀疏化这条路很直观:既然大部分注意力权重接近零,那就干脆不算。代价是你必须提前决定哪些位置可能重要,而这个决定是静态的、与内容无关的。

核化这条路更聪明也更危险:利用 softmax 可以近似分解的性质,把矩阵乘法的顺序换一下,(QK)V(QK^\top)V 变成 Q(KV)Q(K^\top V),复杂度就从 O(n2d)O(n^2 d) 降到 O(nd2)O(n d^2)。代价是这个近似有误差,而误差在长序列上会累积。

各自放弃了什么

我把常见方案按”放弃的能力”整理成一张表。这比按复杂度排序有用得多:

方案复杂度放弃了什么什么任务上免费
滑动窗口O(nw)O(nw)窗口外的精确依赖局部性强的任务,如代码补全
块稀疏O(nn)O(n\sqrt{n})块间细粒度交互长文摘要
LinformerO(n)O(n)序列长度需固定定长输入的分类
PerformerO(n)O(n)精确 softmax,有近似误差对精度不敏感的检索
Mamba / SSMO(n)O(n)任意位置的随机访问流式、因果、长序列
MQA / GQA不变KV 表达多样性推理阶段几乎全免费

最后一行值得单独说:MQA 和 GQA 不降低训练复杂度,它们只减小推理时的 KV cache。这是唯一一个我认为”基本免费”的优化,也是为什么它被普遍采用。

一个被低估的判据

评价这些方案时,大家习惯看困惑度和下游任务分数。但我发现一个更有区分度的测试:在长上下文里插入一条与结尾强相关的信息,看模型能不能用上它。

python
def needle_test(model, ctx_len, depth):
    """把关键信息埋在上下文的 depth 位置,问一个只有用到它才能答对的问题"""
    filler = "无关的背景文字。" * (ctx_len // 8)
    needle = "记住:会议改到了周五下午三点。"

    pos = int(len(filler) * depth)
    prompt = filler[:pos] + needle + filler[pos:] + "\n\n问:会议什么时候?"

    return model.generate(prompt, max_new_tokens=32)

在这个测试上,滑动窗口类方案会在 needle 落在窗口外时直接失败——不是答得差,是完全不知道。而核化方案通常能答出个大概,但细节会漂。这个差异在困惑度上几乎看不出来。

所以该选哪个

我的经验:

  1. 推理慢 → 先上 GQA + PagedAttention,别碰注意力结构。这是性价比最高的一步。
  2. 上下文不够长 → 优先 RoPE 外推或位置插值,比换架构风险小得多。
  3. 真的需要超长序列(十万 token 以上)→ 才考虑 SSM 类方案,并且做好”随机访问能力下降”的预期。

大部分团队的问题在第 1 条就能解决。跳过前两步直接换注意力结构,通常是因为它听起来更像研究工作。

工程上最好的优化,往往是那些不需要重新训练模型的优化。


如果你觉得我某个判断是错的,欢迎写信讨论。