为什么标准注意力是瓶颈?

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等架构的成熟,“注意力机制的平方复杂度"不再是处理长序列的不可逾越的障碍。未来的趋势是混合架构——在不同层使用不同注意力机制,在效率与性能之间取得最优平衡。