Skip to content

长上下文注意力机制论文

扩展语言模型的上下文长度是当前研究的热点方向,从 2K 到 128K 再到百万 Token,长上下文注意力机制的创新使得模型能够处理整本书、长代码库和完整对话历史。

长上下文注意力

核心挑战

标准自注意力的计算复杂度为 O(n²),长序列面临:

  1. 计算瓶颈:序列长度翻倍,注意力计算量翻四倍
  2. 内存瓶颈:KV 缓存随序列长度线性增长
  3. 位置编码外推:训练时未见过的位置导致性能下降

高效注意力机制

稀疏注意力

通过稀疏化注意力矩阵降低计算量:

python
class SlidingWindowAttention(nn.Module):
    """滑动窗口注意力"""
    def __init__(self, d_model, num_heads, window_size=512):
        super().__init__()
        self.num_heads = num_heads
        self.window_size = window_size
        self.qkv = nn.Linear(d_model, 3 * d_model)
        self.out = nn.Linear(d_model, d_model)

    def forward(self, x):
        B, T, D = x.shape
        qkv = self.qkv(x).reshape(B, T, 3, self.num_heads, D // self.num_heads)
        q, k, v = qkv.unbind(dim=2)

        # 创建滑动窗口掩码
        mask = torch.zeros(T, T, dtype=torch.bool, device=x.device)
        for i in range(T):
            start = max(0, i - self.window_size // 2)
            end = min(T, i + self.window_size // 2 + 1)
            mask[i, start:end] = True

        # 仅在窗口内计算注意力
        attn = torch.matmul(q, k.transpose(-2, -1)) / (D // self.num_heads) ** 0.5
        attn = attn.masked_fill(~mask, float('-inf'))
        attn = torch.softmax(attn, dim=-1)
        out = torch.matmul(attn, v)
        return self.out(out.reshape(B, T, D))

FlashAttention

FlashAttention 通过优化 GPU 内存访问实现精确注意力的加速:

  • 分块计算:将注意力矩阵分块,减少 HBM 访问
  • 在线 Softmax:分块内累积 Softmax 统计量
  • 内存节省:无需存储完整 N×N 注意力矩阵

分层注意力

  • 全局+局部:少量 Token 使用全局注意力,其余使用局部窗口
  • 聚类注意力:将相似 Token 聚类,类间计算全局注意力

位置编码扩展

RoPE 扩展

旋转位置编码(RoPE)是当前主流的位置编码方案:

方法策略特点
位置插值缩放位置索引简单但可能损失分辨率
NTK-aware调整 RoPE 基频更好的长度外推
YaRN动态混合缩放结合多种策略
ALiBi线性偏置注意力无需位置编码

NTK-aware 缩放

NTK-aware 缩放的核心洞察:位置编码的高频分量对局部关系敏感,低频分量对全局关系敏感。扩展上下文时,应调整 RoPE 的基频 θ,使得新的频率分布能覆盖更长的序列范围。

KV 缓存优化

长上下文的 KV 缓存是内存瓶颈:

  • PagedAttention:虚拟内存式的 KV 缓存管理
  • KV 压缩:注意力加权合并相似 KV 对
  • StreamingLLM:保留开头和最近的 KV,丢弃中间
  • 量化 KV:将 KV 缓存量化到更低精度

相关资源

最近更新