为什么标准注意力是瓶颈?
Transformer的核心——自注意力机制,其计算复杂度是序列长度的二次方 O(n²)。这个n²使得处理长序列时的计算和内存开销急剧增长。
标准注意力的计算过程
import torch
import torch.nn.functional as F
import math
def standard_attention(Q, K, V):
"""
标准自注意力
Q, K, V: (batch, num_heads, seq_len, d_head)
"""
d = Q.shape[-1]
# 注意力分数: (batch, heads, seq, seq)
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d)
# Softmax归一化
attn = F.softmax(scores, dim=-1)
# 加权求和
output = torch.matmul(attn, V) # (batch, heads, seq, d_head)
return output
# 计算复杂度分析
# scores矩阵大小: seq_len × seq_len
# 当 seq_len = 8192 时, 每个head的scores矩阵 = 8192² = 67M 个元素
# 当 seq_len = 32768 时, 每个head的scores矩阵 = 32768² = 1B 个元素
# 显存随序列长度平方增长!
不同序列长度的开销对比
| 序列长度 | 注意力矩阵大小 | 显存(FP16, 32头) | 计算量(FLOPs) |
|---|---|---|---|
| 4K | 16M | 1GB | 0.13T |
| 8K | 64M | 4GB | 0.52T |
| 32K | 1B | 64GB | 8.4T |
| 128K | 16B | 1TB | 134T |
| 1M | 1T | 64TB | 8.2PT |
这个表格清楚地展示了为什么标准注意力无法扩展到超长序列。
线性注意力的核心思想
线性注意力的关键洞察是:可以通过改变计算顺序将O(n²)降为O(n)。
数学推导
标准注意力可以写为:
$$\text{Attn}(Q,K,V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d}}\right)V$$
如果将softmax替换为任意核函数 $\phi(\cdot)$,则:
$$\text{Attn}_{\text{linear}}(Q,K,V) = \phi(Q)(\phi(K)^T V)$$
关键在于计算顺序:
- 标准顺序:$(Q K^T) V$ → 先算 $n \times n$ 矩阵
- 线性顺序:$Q (K^T V)$ → 先算 $d \times d$ 矩阵,与 $n$ 无关!
def linear_attention(Q, K, V, feature_map=None):
"""
线性注意力
通过改变计算顺序将复杂度从O(n²d)降为O(nd²)
当 d << n 时,显著降低计算量
"""
if feature_map is None:
# ELU+1 特征映射
feature_map = lambda x: F.elu(x) + 1
Q_prime = feature_map(Q) # (batch, heads, seq, d)
K_prime = feature_map(K) # (batch, heads, seq, d)
# 关键:先计算 K^T V (d × d),而非 Q K^T (n × n)
KV = torch.matmul(K_prime.transpose(-2, -1), V) # (batch, heads, d, d)
# 再用 Q 乘以 KV
output = torch.matmul(Q_prime, KV) # (batch, heads, seq, d)
# 归一化
K_sum = K_prime.sum(dim=-2, keepdim=True) # (batch, heads, 1, d)
normalizer = torch.matmul(Q_prime, K_sum.transpose(-2, -1)) # (batch, heads, seq, 1)
output = output / (normalizer + 1e-6)
return output
计算复杂度对比
| 方案 | 时间复杂度 | 空间复杂度 | 支持因果掩码 |
|---|---|---|---|
| 标准注意力 | O(n²d) | O(n²) | ✅ 直接支持 |
| 线性注意力 | O(nd²) | O(nd) | ⚠️ 需要特殊处理 |
| Flash Attention | O(n²d) | O(n) | ✅ |
| Mamba/SSM | O(nd) | O(d) | ✅ 原生支持 |
核心架构解析
1. Linear Transformer (2020)
Katharopoulos等人提出的线性注意力是最基础的方案:
特征映射选择:
- $\phi(x) = \text{ELU}(x) + 1$:简单高效,但表达能力有限
- $\phi(x) = \text{ReLU}(x)$:更简单,实践中效果接近
- 随机特征映射:近似softmax更精确,但增加计算量
因果掩码的挑战:线性注意力中,因果掩码需要在每个时间步维护一个累积状态:
class CausalLinearAttention:
"""支持因果掩码的线性注意力(递归形式)"""
def __init__(self, d_head):
self.d = d_head
# 累积状态 S_t = sum_{i<=t} phi(K_i) V_i^T
self.S = torch.zeros(d_head, d_head)
# 累积归一化因子
self.z = torch.zeros(d_head)
def step(self, q, k, v):
"""逐时间步更新"""
q_prime = F.elu(q) + 1
k_prime = F.elu(k) + 1
# 使用累积状态计算输出
output = torch.matmul(q_prime, self.S) / (torch.dot(q_prime, self.z) + 1e-6)
# 更新状态
self.S += torch.outer(k_prime, v)
self.z += k_prime
return output
这种递归形式使得Linear Transformer可以像RNN一样处理无限长度序列。
2. RWKV (2023-2024)
RWKV是由BlinkDL社区开发的混合架构,结合了RNN的推理效率与Transformer的表达能力:
核心设计:
- WKV机制:替代标准注意力的时间混合方式
- Token Shift:相邻Token之间的信息传递
- Channel Mix:替代FFN的非线性变换
class RWKVTimeMixing(nn.Module):
"""RWKV的时间混合层核心逻辑"""
def __init__(self, d_model):
super().__init__()
self.time_mix_k = nn.Parameter(torch.ones(1) * 0.5)
self.time_mix_v = nn.Parameter(torch.ones(1) * 0.5)
self.time_mix_r = nn.Parameter(torch.ones(1) * 0.5)
self.key = nn.Linear(d_model, d_model, bias=False)
self.value = nn.Linear(d_model, d_model, bias=False)
self.receptance = nn.Linear(d_model, d_model, bias=False)
self.output = nn.Linear(d_model, d_model, bias=False)
def forward(self, x, state):
"""
x: (batch, d_model) 当前时间步输入
state: (batch, d_model) 上一时间步的状态
"""
# Token Shift: 当前输入与历史状态的线性混合
xk = x * self.time_mix_k + state * (1 - self.time_mix_k)
xv = x * self.time_mix_v + state * (1 - self.time_mix_v)
xr = x * self.time_mix_r + state * (1 - self.time_mix_r)
k = self.key(xk)
v = self.value(xv)
r = torch.sigmoid(self.receptance(xr))
# WKV计算(简化版)
w = torch.exp(-torch.exp(k)) # 衰减因子
output = r * self.output(v)
return output, x # 返回输出和新状态
RWKV的优势:
- 推理时为O(1)空间复杂度(只需维护固定大小状态)
- 训练时可并行(类似Transformer)
- 在7B-14B参数规模下性能接近同规模Transformer
3. Mamba (2023-2025)
Mamba基于状态空间模型(SSM),通过选择性机制实现了对序列内容的自适应处理:
SSM基础:
状态空间模型可以表示为连续动态系统:
$$h’(t) = Ah(t) + Bx(t)$$ $$y(t) = Ch(t) + Dx(t)$$
离散化后:
$$h_t = \bar{A}h_{t-1} + \bar{B}x_t$$ $$y_t = Ch_t + Dx_t$$
Mamba的关键创新——选择性SSM:
class MambaBlock(nn.Module):
"""Mamba选择性状态空间模型核心逻辑"""
def __init__(self, d_model, d_state=16, d_conv=4, expand=2):
super().__init__()
d_inner = d_model * expand
# 输入投影
self.in_proj = nn.Linear(d_model, d_inner * 2, bias=False)
# 1D卷积(局部特征提取)
self.conv = nn.Conv1d(d_inner, d_inner, d_conv,
padding=d_conv-1, groups=d_inner)
# SSM参数投影(选择性机制的核心)
self.x_proj = nn.Linear(d_inner, d_state * 2 + d_inner, bias=False)
# A参数(通过Log参数化保证负定性)
A = torch.arange(1, d_state + 1).float().repeat(d_inner, 1)
self.A_log = nn.Parameter(torch.log(A))
self.D = nn.Parameter(torch.ones(d_inner))
# 输出投影
self.out_proj = nn.Linear(d_inner, d_model, bias=False)
def ssm_step(self, x_t, h_prev, A_bar, B_bar, C):
"""单步SSM计算"""
# 状态更新: h_t = A_bar * h_{t-1} + B_bar * x_t
h_t = A_bar * h_prev + B_bar * x_t
# 输出: y_t = C * h_t + D * x_t
y_t = torch.einsum('bd,bsd->bs', C, h_t) + self.D * x_t
return y_t, h_t
def forward(self, x):
"""
x: (batch, seq_len, d_model)
"""
batch, seq_len, _ = x.shape
# 投影和分割
xz = self.in_proj(x) # (batch, seq, 2*d_inner)
x_branch, z_branch = xz.chunk(2, dim=-1)
# 卷积
x_conv = self.conv(x_branch.transpose(1, 2)).transpose(1, 2)
x_conv = F.silu(x_conv)[:, :seq_len, :]
# 选择性参数(根据输入动态计算A, B, C)
ssm_params = self.x_proj(x_conv) # (batch, seq, d_state*2 + d_inner)
B, C, delta = ssm_params.split(
[self.d_state, self.d_state, x_branch.shape[-1]], dim=-1
)
# 离散化步长(选择性机制:delta依赖于输入)
delta = F.softplus(delta) # 保证正定
# A的离散化: A_bar = exp(delta * A)
A = -torch.exp(self.A_log) # 保证负定
A_bar = torch.exp(delta.unsqueeze(-1) * A.unsqueeze(0).unsqueeze(0))
# B的离散化: B_bar = delta * B
B_bar = delta.unsqueeze(-1) * B.unsqueeze(1)
# 递归计算(实际实现中使用并行扫描算法)
h = torch.zeros(batch, x_branch.shape[-1], self.d_state, device=x.device)
outputs = []
for t in range(seq_len):
y_t, h = self.ssm_step(x_conv[:, t, :], h,
A_bar[:, t], B_bar[:, t], C[:, t])
outputs.append(y_t)
y = torch.stack(outputs, dim=1)
# 门控
y = y * F.silu(z_branch)
return self.out_proj(y)
Mamba vs Transformer vs Linear Attention对比:
| 特性 | Transformer | Linear Attn | RWKV | Mamba |
|---|---|---|---|---|
| 训练并行度 | ✅ 完全并行 | ✅ 完全并行 | ✅ 并行 | ✅ 并行(扫描) |
| 推理复杂度 | O(n²d) | O(nd²) | O(d) | O(d) |
| 长程依赖 | ✅ 强 | ⚠️ 弱 | ⚠️ 中等 | ✅ 强 |
| 内容自适应 | ✅ 是 | ❌ 否 | ❌ 否 | ✅ 是 |
| 硬件效率 | 高 | 中等 | 高 | 高 |
工程实践建议
何时选择线性注意力?
选择标准Transformer:
- 序列长度 < 8K
- 需要精确的长程注意力
- 训练与推理预算充足
选择Linear Attention/RWKV:
- 需要流式推理(无限长度输入)
- 边缘设备部署(内存受限)
- 可以接受轻微性能损失
选择Mamba:
- 长序列处理(>32K)
- 需要内容自适应的序列建模
- 推理延迟敏感场景
混合架构(当前最优方案):
- 底层使用Mamba/线性注意力处理长序列
- 顶层穿插少量标准注意力层捕获全局依赖
- 如Jamba、Zamba等混合架构
性能基准(2026年)
| 架构 | 模型规模 | MMLU | 长文本理解 | 推理速度(相对) |
|---|---|---|---|---|
| Transformer | 7B | 65.2 | 强 | 1.0x |
| RWKV-7 | 7B | 63.8 | 中等 | 3.2x |
| Mamba-2 | 7B | 64.5 | 强 | 2.8x |
| Jamba (混合) | 12B/3B(MoE) | 68.1 | 强 | 2.1x |
线性注意力及其衍生架构已经从学术探索走向工程落地。随着Mamba、RWKV等架构的成熟,“注意力机制的平方复杂度"不再是处理长序列的不可逾越的障碍。未来的趋势是混合架构——在不同层使用不同注意力机制,在效率与性能之间取得最优平衡。