长上下文与位置编码

当大模型的上下文窗口从 2K、4K 一路扩展到 32K、128K 甚至更长,两个底层问题始终绕不开:模型如何知道词与词之间的先后顺序,以及如何在可接受的计算与显存成本下处理成千上万个 token。前者是位置编码(Positional Encoding)要解决的,后者则催生了 FlashAttention、Longformer 等一长串高效注意力方案。本文从「为什么需要位置」讲起,依次拆解绝对位置编码、相对位置编码、RoPE、ALiBi,再回到长上下文工程与 lost in the middle 现象。

为什么自注意力需要位置信息

标准自注意力的计算为:

Attention(Q, K, V) = softmax(Q · Kᵀ / sqrt(d)) · V

其中 Q、K、V 由输入向量线性投影得到。关键在于:这个等式对输入序列的「顺序」是完全无关的。如果你把输入的第 i 个和第 j 个 token 交换,输出的第 i 个和第 j 个向量也会原样交换——数学上称这种性质为置换等变(permutation equivariant)。换句话说,自注意力本身是「无序」的,它分不清「猫追狗」和「狗追猫」。

为了理解语言,模型必须知道词序。因此必须在某处注入位置信息。注入方式大致分两类:

  • 把位置信号「加」进输入向量(绝对位置编码的代表);
  • 把位置信号「融」进注意力分数的计算(相对位置编码、RoPE、ALiBi 的代表)。

绝对位置编码:正弦/余弦函数

Vaswani 等人 2017 年在原始 Transformer 中提出用一组固定的正弦、余弦函数给每个位置生成一个向量,再直接加到词嵌入上:

PE(pos, 2i)     = sin( pos / 10000^(2i/d) )
PE(pos, 2i+1)   = cos( pos / 10000^(2i/d) )

其中 pos 是位置下标,i 是维度下标,d 是模型维度。每一个维度对应一条频率不同的正弦波,低频维度编码「大概位置」、高频维度编码「精细位置」。这种编码有两条重要性质:

  • 不同位置对应的向量彼此不同,可被模型区分;
  • 任意固定偏移 k,PE(pos+k) 都可以由 PE(pos) 通过线性变换得到,这暗示模型有可能据此学到相对位置规律。

下面是生成该编码的参考实现:

import torch

def sinusoidal_positional_encoding(seq_len, d_model):
    # position: (seq_len, 1)
    position = torch.arange(seq_len).unsqueeze(1).float()
    # div_term = 10000^(-2i/d),对应 1 / 10000^(2i/d)
    div_term = torch.exp(
        torch.arange(0, d_model, 2).float()
        * (-torch.log(torch.tensor(10000.0)) / d_model)
    )
    pe = torch.zeros(seq_len, d_model)
    pe[:, 0::2] = torch.sin(position * div_term)
    pe[:, 1::2] = torch.cos(position * div_term)
    return pe  # 形状 (seq_len, d_model)

绝对位置编码简单、无需训练,但有两个局限:位置是「绝对」的,而语言中的依赖更依赖「相对距离」;且训练时见过的最大长度之外,泛化能力较弱。

相对位置编码

语言里的很多关系由相对距离决定:动词和就近宾语的依存,并不取决于它们分别在第几个位置,而取决于它们之间隔了多远。相对位置编码正是把「相对距离」直接编码进注意力。两条有影响力的路线如下。

Transformer-XL 的相对位置编码

Dai 等人 2019 年在 Transformer-XL 中把注意力分数重写成四项之和,使位置信息以「相对距离」的形式进入计算。其核心是:键向量被拆成「内容部分」和「位置部分」,而位置部分只依赖相对距离 i−j,而不是绝对位置。具体把分数拆为:

a(i, j) = 内容-内容 + 内容-相对位置 + 绝对位置-内容 + 相对位置-相对位置

其中相对位置用一组可学习的正弦向量 R(i−j) 表示。这样模型在做长距离建模(并配合段循环缓存)时,学到的是「相距多少」而非「在第几」。

T5 的相对位置偏置

Raffel 等人 2020 年的 T5 干脆不再把任何位置向量加到词嵌入上,而是在自注意力里给 query-key 点积加一个「相对位置偏置」。对于位置 j 与 k 的相对距离,T5 先对距离做分桶(bucketing):近距离用细粒度单桶,远距离用对数间隔合并到同一桶,桶内共享同一个可学习标量参数。偏置项直接加到注意力分数上:

score(j, k) = q_jᵀ k_k + bias(bucket(|j - k|))

分桶让 T5 既能表达相对距离,又能把很长的距离压缩进有限桶数,从而获得较好的长度外推性。

RoPE 旋转位置编码

RoPE(Rotary Position Embedding,旋转位置编码)由 Su 等人于 2021 年提出(arXiv:2104.09864,对应模型 RoFormer),是当下 Llama、Qwen、ChatGLM 等主流大模型采用的位置方案。它的设计目标是:在编码「绝对位置」的同时,让注意力分数天然只依赖于「相对距离」。

直觉

把每个位置的查询/键向量看作一组二维平面上的向量,按位置下标做「旋转」:位置 m 的向量整体旋转 m·θ 角度。两个向量各自旋转后再做点积,旋转角度会相互抵消成「相对角度」,于是点积只与它们的相对距离有关。绝对位置被编码进了旋转,相对位置则在注意力分数里自然浮现。

二维情形

把二维向量 (x, y) 旋转角度 φ 写为矩阵乘法:

R(φ) = [ cos φ, -sin φ ]
       [ sin φ,  cos φ ]

设位置 m 的键向量为 k_m = R(mθ)·k,位置 n 的查询向量为 q_n = R(nθ)·q,则它们的点积满足:

q_nᵀ k_m = qᵀ R((m - n)θ) k

右侧只含相对距离 (m−n),这正是「相对位置依赖」的来源。

高维推广

对 d 维向量,把维度两两配对 (x_{2i}, x_{2i+1}),第 i 对按角度 m·θ_i 旋转,其中 θ_i = 10000^(-2i/d)。旋转后的查询在第 i 对上的分量为:

q_m,[2i]   = q_{2i}   cos(m θ_i) - q_{2i+1} sin(m θ_i)
q_m,[2i+1] = q_{2i}   sin(m θ_i) + q_{2i+1} cos(m θ_i)

最终查询与键的点积可以化简为只含 (m−n) 的形式:

q_mᵀ k_n = Σ_i [ (q_{2i} k_{2i} + q_{2i+1} k_{2i+1}) cos((m-n)θ_i)
               + (q_{2i} k_{2i+1} - q_{2i+1} k_{2i})   sin((m-n)θ_i) ]

RoPE 由此具备几项好性质:相对位置被显式编码进注意力;对相对距离有随距离增大而衰减的依赖;序列长度灵活。它也带来一个工程现实——直接把训练长度外推到远超训练范围的 m 会退化,因此社区发展出 NTK-aware、线性缩放(linear scaling)等 RoPE 缩放技巧来延长上下文。

下面是 RoPE 的常用实现骨架:

import torch

def precompute_rope_freqs(dim: int, base: float = 10000.0):
    # dim 为注意力头维度,需为偶数;返回 (dim/2,) 的频率项
    half = dim // 2
    inv_freq = 1.0 / (base ** (torch.arange(0, half, dtype=torch.float32) * 2.0 / dim))
    return inv_freq

def apply_rope(x: torch.Tensor, inv_freq: torch.Tensor):
    # x: (batch, seq_len, n_head, head_dim)
    seq_len = x.size(1)
    pos = torch.arange(seq_len, dtype=torch.float32)
    angles = torch.outer(pos, inv_freq)      # (seq_len, half)
    cos = torch.cos(angles)                  # (seq_len, half)
    sin = torch.sin(angles)                  # (seq_len, half)
    x_even = x[..., 0::2]                     # 每对的第一维
    x_odd = x[..., 1::2]                      # 每对的第二维
    rot_even = x_even * cos - x_odd * sin
    rot_odd = x_even * sin + x_odd * cos
    out = torch.zeros_like(x)
    out[..., 0::2] = rot_even
    out[..., 1::2] = rot_odd
    return out  # 旋转后的查询/键

ALiBi 线性偏置

ALiBi(Attention with Linear Biases,带线性偏置的注意力)由 Press 等人于 2021 年提出(arXiv:2108.12409),标题直白地说清了它的卖点:训练短、测试长(Train Short, Test Long)。

ALiBi 的思路与前面都不同:它完全不往词嵌入上加任何位置向量,而是直接给 query-key 注意力分数加一个「随距离线性增长」的惩罚:

score(i, j) = q_iᵀ k_j  -  m_h · |i - j|

其中 |i−j| 是位置 i 与 j 的距离,m_h 是每个注意力头各自的斜率。距离越远,惩罚越大,分数被压得越低;不同头用不同 m_h,使模型同时关注近处与远处(典型做法是让各头斜率组成几何级数,靠前的头斜率大、偏重近期)。

这种「偏好近期」的归纳偏置带来了关键能力:在长度 1024 上训练,也能外推到 2048 做推理,且比正弦位置编码训得更快(论文报告训练快约 11%、省约 11% 显存)。下面是构造 ALiBi 偏置矩阵的参考代码:

import torch

def alibi_bias(seq_len: int, n_heads: int, device="cpu"):
    # 每个头一个斜率 m_h,构成几何级数(论文以 8 头为例,斜率呈 2 的负幂次)
    # 此处斜率公式仅作示意,核心形式是 bias = -m_h * |i - j|
    slopes = torch.tensor(
        [2.0 ** (-(8.0 * (h + 1) / n_heads)) for h in range(n_heads)]
    )
    dist = torch.arange(seq_len, device=device).unsqueeze(0) \
        - torch.arange(seq_len, device=device).unsqueeze(1)
    dist = dist.abs().unsqueeze(0)                       # (1, seq_len, seq_len)
    bias = -slopes.view(-1, 1, 1) * dist                 # (n_heads, seq_len, seq_len)
    return bias

高效长上下文注意力

位置编码解决「顺序」,但自注意力的计算复杂度是序列长度 N 的平方(O(N²)),长上下文的真正瓶颈在算力与显存。两条互补路线:把精确注意力做得更省(FlashAttention),或把注意力变稀疏(Longformer)。

FlashAttention

Dao 等人 2022 年提出 FlashAttention(arXiv:2205.14135)。标准注意力要把 N×N 的分数矩阵完整写进 GPU 高带宽显存(HBM),反复读写导致又慢又占显存。FlashAttention 的核心观察是「IO 感知」:把 Q、K、V 切成能放进片上 SRAM 的小块,用在线 softmax(online softmax)在分块流式计算中直接得到最终结果,全程不物化完整的 N×N 矩阵。

结果是:注意力仍然是精确的(不是近似),但显存降到 O(N),且 wall-clock 明显加速。论文还把它扩展到分块稀疏注意力。FlashAttention 让 Transformer 能吃下更长上下文(论文中首次在 16K 长度的 Path-X、64K 长度的 Path-256 上取得优于随机的表现)。

Longformer

Beltagy 等人 2020 年提出 Longformer(arXiv:2004.05150),用稀疏注意力把复杂度降到随序列长度线性。它把标准全注意力替换成两种注意力的组合:

  • 滑动窗口局部注意力:每个 token 只关注左右各 w 个邻域 token,复杂度 O(N·w);
  • 全局注意力:一小撮「任务指定的」token(如分类的 CLS、问答的问题 token)对所有 token 可见、也被所有 token 关注。

两者结合既保留了长程关键信息通道,又把计算压到线性。Longformer 因此能直接处理数千 token 的文档,并在 WikiHop、TriviaQA 等长文档任务上优于 RoBERTa。

同方向的还有 Sparse Transformer、Linformer、BigBird 等:思路要么对注意力矩阵做低秩/稀疏近似,要么限制每个 token 的关注范围。FlashAttention 的特别之处在于它「不近似」也省资源。

Lost in the Middle 问题

即便模型声称支持很长的上下文窗口,也未必「用得好」。Liu 等人 2023 年在「Lost in the Middle」(arXiv:2307.03172)中做了系统实验:在多文档问答与键值检索两类任务上,当关键信息出现在输入的开头或结尾时,模型表现最好;一旦关键信息被放到长上下文的「中段」,表现会明显下滑——即便模型本身标榜支持长上下文。

这给工程上的启示是:单纯把上下文窗口拉大,不等于模型能稳定利用中间的信息。实践中可以:

  • 把最关键的内容放在提示的开头或结尾;
  • 对超长材料做分块检索(RAG),而不是整篇塞进上下文;
  • 评估长上下文能力时,要在「信息位于不同位置」的条件下分别测,而非只测首尾。

关于长上下文评测与检索增强的更多实践,站内另有专文,可配合本文阅读。

小结

  • 自注意力本身置换等变、不分词序,必须注入位置信息才能理解语言。
  • 绝对位置编码(正弦/余弦)简单但偏「绝对」,外推能力有限。
  • 相对位置编码(Transformer-XL、T5)把位置以相对距离形式融进注意力,更贴合语言依赖。
  • RoPE 用旋转把绝对位置编码进向量,使注意力分数天然只依赖相对距离,被 Llama/Qwen/ChatGLM 等广泛采用。
  • ALiBi 不加位置向量,而是给注意力分数加随距离线性增长的惩罚,能「训练短、测试长」地外推。
  • FlashAttention 用 IO 感知的分块计算做精确注意力,把显存降到线性、显著加速;Longformer 用局部窗口加全局注意力实现线性复杂度。
  • lost in the middle 提醒我们:长上下文窗口不等于模型善用中段信息,关键内容应放在首尾或配合检索。

参考与延伸阅读

本文累计阅读