超越自注意力:Transformer 之后是什么
Transformer 的二次方注意力开销是真实瓶颈。线性注意力、注意力残差与混合架构,指向了下一代模型的方向。

Transformer 架构已经支撑 AI 领域八年了。如今主流的大语言模型、大多数图像生成系统,以及越来越多的音频和视频模型,都建立在《Attention Is All You Need》论文提出的自注意力机制之上。但自注意力有个根本性的问题:它的计算和内存开销随序列长度呈二次方增长。输入长度翻一倍,开销就翻四倍。
序列短的时候,这个问题不明显。但对于我们正在追求的 128K token 上下文窗口,以及大家都想要的百万 token 窗口来说,它就是实打实的瓶颈。现在有一批研究在探索替代方案:跨层复用计算的注意力残差、摆脱二次方开销的线性注意力变体,以及把注意力与更廉价机制混合起来的混合架构。Transformer 不会消失,但它正在被重塑。
为什么自注意力这么贵
要理解这些替代方案,先得搞清楚自注意力到底在算什么。给定一个包含 N 个 token 的序列,自注意力会计算每两个 token 之间的相关性分数:token 1 和 token 2、token 1 和 token 3、一直到 token 1 和 token N,然后是 token 2 和其他所有 token,依此类推。总共是 N² 对。
import torch
import torch.nn.functional as F
def self_attention(Q, K, V):
"""
Standard self-attention.
Q, K, V: (batch, seq_len, d_model)
The attention matrix is seq_len × seq_len.
For seq_len = 1024: ~1M entries (manageable)
For seq_len = 32768: ~1B entries (expensive)
For seq_len = 131072: ~17B entries (very expensive)
"""
d_k = Q.size(-1)
# This matmul creates the N×N attention matrix
scores = torch.matmul(Q, K.transpose(-2, -1)) / (d_k ** 0.5)
weights = F.softmax(scores, dim=-1)
return torch.matmul(weights, V)
在 4K token 下,注意力矩阵有 1600 万个元素,现代 GPU 完全扛得住。到了 128K token,就有 160 亿 个元素。到了 100 万 token,则超过一万亿。即便用了 Flash Attention(它并没有减少计算量,但极大改善了内存访问模式),二次方的增长最终还是会占上风。
这也是为什么早期的 Transformer 只能处理 512 或 1024 个 token。每一代硬件和优化手段都把上限往上推了一些,但我们撞上的是一堵数学上的墙。线性增长(O(N))从根本上优于二次方增长(O(N²)),这正是大多数替代架构追求的目标。
注意力残差:复用已经算过的结果
降低注意力开销最务实的方法之一,并不是替换注意力,而是通过复用前面层的计算结果,让每一层注意力变得更便宜。
观察结果是:在一个深层 Transformer(比如 32 层)中,相邻层的注意力模式往往惊人地相似。第 15 层和第 16 层通常关注的位置差不多,只有细微调整。每一层都从头计算完整的 N² 注意力矩阵,其实是一种冗余,因为其中大部分工作在前一层就已经做过了。
注意力残差正是利用了这一点:它计算的是一种“残差”注意力模式,即当前层想关注的内容与前一层已经算出的结果之间的差值。如果差值很小(中间层通常如此),计算就会便宜得多。完整的注意力模式等于前一层的模式加上当前层的残差。
这和视频压缩的思路很像:与其独立存储每一帧,不如存一个关键帧,再存一系列相对于关键帧的差分(残差)。差分通常比完整帧小得多,所以压缩效果会好很多。
实践中,在深层模型的中间层,注意力残差能把注意力的计算开销降低 30% 到 50%,而对质量的影响很小。前几层和最后几层仍然需要完整的注意力计算(它们的模式更加独特),但占大多数的中间层能获得显著的加速。
线性注意力:摆脱二次方开销
线性注意力变体试图改写注意力的计算方式,使其复杂度从 O(N²) 降到 O(N)。大体思路是:不再显式计算 N×N 的注意力矩阵,而是通过线性运算得到相同(或近似相同)的输出。
其中的数学技巧依赖于 softmax 的核函数分解。标准注意力计算的是 softmax(QK^T)V。如果用一个可以分解为 φ(Q) · φ(K)^T 的其他核函数替换 softmax,就能调整计算顺序:不再计算 (φ(Q) · φ(K)^T) · V(中间结果是 N×N),而是计算 φ(Q) · (φ(K)^T · V)(中间结果是 d×d,其中 d 是模型维度)。对于长序列来说 d 远小于 N,因此这种计算方式便宜得多。
def linear_attention(Q, K, V, feature_map=None):
"""
Linear attention via kernel feature maps.
Cost: O(N * d^2) instead of O(N^2 * d)
"""
if feature_map is None:
# ELU+1 is a common choice (from Katharopoulos et al.)
feature_map = lambda x: F.elu(x) + 1
Q = feature_map(Q) # (batch, seq_len, d)
K = feature_map(K) # (batch, seq_len, d)
# Key insight: compute K^T @ V first (d × d matrix)
# instead of Q @ K^T first (N × N matrix)
KV = torch.einsum('bnd,bnm->bdm', K, V) # (batch, d, d)
# Then multiply by Q
output = torch.einsum('bnd,bdm->bnm', Q, KV) # (batch, N, d)
# Normalize
Z = torch.einsum('bnd,bd->bn', Q, K.sum(dim=1)) # normalization
output = output / Z.unsqueeze(-1)
return output
问题在于:用别的核函数替换 softmax 会改变注意力的分布,而用 softmax 训练出来的模型,不一定能很好地迁移到线性注意力上。近来的线性注意力变体已经把质量差距缩小了很多,能达到 softmax 注意力 95% 到 98% 的质量,但差距仍然存在,尤其是在需要精确长距离检索的任务上。
状态空间模型:另一种范式
像 Mamba 这样的状态空间模型(SSM)走的是完全不同的路线。它们不计算 token 之间的两两关系,而是通过递归来处理序列,维护一个固定大小的隐状态,并在每个 token 处更新。这种方式天然是 O(N):处理的 token 数量翻倍,耗时也只翻倍,而不是翻四倍。
现代 SSM 的创新之处在于,让递归的参数依赖于输入(即选择性状态空间)。这让模型具备了一种基于内容的注意力:它可以“选择”记住什么、遗忘什么,而且没有二次方开销。类 Mamba 的模型在很多基准上能与 Transformer 持平,同时在长序列上明显更快。
权衡在于:SSM 需要按顺序处理 token,因此相比 Transformer(可以同时处理所有 token),训练时更难并行化。训练效率很重要:一个推理快 2 倍、训练却慢 3 倍的模型不一定划算,毕竟总计算量的大头都花在训练上。
混合架构:务实的路线
目前生产环境模型的趋势,是组合不同注意力机制的混合架构。道理很简单:模型的不同部分适合不同类型的计算。
- 用完整注意力做全局推理。有些层需要横跨整个序列,找到相隔数千个 token 之外的相关上下文。这些层使用标准(可能经过 Flash 优化)的自注意力。
- 用局部注意力处理邻近上下文。许多层主要关注附近的 token(滑动窗口注意力)。使用 256 到 1024 个 token 的固定窗口,可以把开销降到 O(N·W),其中 W 是窗口大小。
- 用线性注意力聚合广域上下文。有些层需要汇总整个序列的信息,但并不需要精确的注意力权重。线性注意力以 O(N) 的成本就能做到这一点。
- 用 SSM 层处理顺序依赖。类 Mamba 的层无需任何注意力计算,就能高效处理序列中的依赖关系。
Jamba(AI21)等模型以及各种研究架构,会根据层的角色在这些机制之间交替使用。前几层使用局部注意力(处理语法和局部模式)。中间层使用线性注意力或 SSM(构建更广的表示)。少数关键层使用完整注意力(负责全局推理和检索)。这样既能实现接近线性的整体扩展性,又保留了那些确实需要完整注意力的模型质量。
开发者应该关注什么
如果你正在基于大语言模型开发应用,底层架构的这些变化会以具体的方式影响你的工作。
- 上下文窗口会持续增长。随着注意力成本下降,上下文窗口也会随之扩大。这会改变应用的架构:过去需要搭建复杂的 RAG 管道,把相关内容塞进 4K 窗口里;现在你或许可以直接把所有内容塞进一个 100 万 token 的提示词里。这种简单性很有吸引力,但不同架构在延迟和成本上的表现差别很大。
- 延迟特征会发生变化。Transformer 在一定长度之内延迟相对平稳,超过之后则呈二次方上升。线性注意力和 SSM 模型的延迟则是更平缓的线性增长。对于响应时间很关键的应用,了解你所用模型的扩展行为非常重要。
- 质量差异取决于任务。线性注意力模型在需要从长上下文特定位置精确检索信息的任务上(比如“第 47 页列表里的第三项是什么?”)可能略逊一筹,但在需要整体理解的任务上表现同样出色。关键是要清楚你的使用场景。
- 推理优化变得更加重要。随着模型的架构越来越复杂(混合不同类型的注意力),推理引擎需要高效处理异构计算。vLLM、TensorRT-LLM 等框架正在适配,但自定义架构可能不会立即获得支持。
Transformer 并不是被取代,而是在不断演进。自注意力依然是我们手里对 token 之间关系建模表达能力最强的机制。但它不必在每一层、每个位置都以完整的 N² 成本使用。未来几年的模型会精准地使用注意力:在最关键的地方保持完整精度,其余地方则采用更廉价的替代方案。结果就是模型更快、能处理更长的上下文、运行成本更低,同时质量持平甚至超越现在。这值得我们持续关注。


