为什么二维旋转解决了 Transformer 长上下文:RoPE 的机制优先视角
注意力机制对位置编码有两项要求:
- 相对距离敏感性:第 ii 个 token 关注第 jj 个 token 时,主要应关心两者之间的距离( i−ji - j ),而非它们在上下文窗口中的全局位置。
- 特征保持性:注入位置信息时不得破坏或扭曲模型学到的语义嵌入特征。
原始绝对位置编码将静态正弦/余弦波直接加到输入嵌入上:
xi=ei+pi \mathbf{x}_i = \mathbf{e}_i + \mathbf{p}_i
这迫使模型消耗参数容量来“解混”语义含义( ei\mathbf{e}i )与位置信息( pi\mathbf{p}_i )。更糟糕的是,若模型训练时最大序列长度为 N=2048N = 2048 ,位置 20492049 对应的位置向量 p2049\mathbf{p}{2049} 属于未见过的向量,会导致生成立即崩溃。
后续方法尝试将相对偏置项 bi,jb_{i,j} 直接加到注意力矩阵 QKT+BQK^T + B 中。虽然在功能上有效,但修改 N×NN \times N 矩阵会破坏 FlashAttention 等 GPU 内核融合,并引入巨大内存开销。
于是出现了 旋转位置编码(RoPE)(Su et al., 2021)。RoPE 既不 添加 位置向量,也不 修补 注意力矩阵,而是先在 2D 子平面上 旋转 Query 和 Key 向量,再计算注意力。
几何视角:旋转下的内积
为使注意力具有相对性,我们希望一个编码函数 R(x,m)R(\mathbf{x}, m) 作用于位置 mm 处的向量 x\mathbf{x} ,使得位置 mm 的 Query 与位置 nn 的 Key 之间的点积 仅取决于相对偏移( m−nm - n ):
⟨R(q,m),R(k,n)⟩=g(q,k,m−n) \langle R(\mathbf{q}, m), R(\mathbf{k}, n) \rangle = g(\mathbf{q}, \mathbf{k}, m - n)
如何在编码角度偏移的同时保持向量范数?答案是 复空间旋转。
在 2D 平面上,将向量 x=[x1,x2]T\mathbf{x} = [x_1, x_2]^T 旋转 mθm\theta 角度可用正交旋转矩阵表示:
RΘ,m2=(cosmθ−sinmθ sinmθcosmθ) R_{\Theta, m}^2 = \begin{pmatrix} \cos m\theta & -\sin m\theta \ \sin m\theta & \cos m\theta \end{pmatrix}
计算位置 mm 的旋转后 Query 与位置 nn 的旋转后 Key 的点积:
(RΘ,m2q)T(RΘ,n2k)=qT(RΘ,m2)TRΘ,n2k=qTRΘ,n−m2k \left(R_{\Theta, m}^2 \mathbf{q}\right)^T \left(R_{\Theta, n}^2 \mathbf{k}\right) = \mathbf{q}^T \left(R_{\Theta, m}^2\right)^T R_{\Theta, n}^2 \mathbf{k} = \mathbf{q}^T R_{\Theta, n - m}^2 \mathbf{k}
由于 RT(m)R(n)=R(n−m)R^T(m) R(n) = R(n - m) ,绝对位置 mm 和 nn 完全被消去。最终的注意力权重仅是距离( m−nm - n )的函数。
对于 dd 维向量,RoPE 将通道分成 d/2d/2 个 2D 平面对,每对使用不同的旋转频率 θi\theta_i :
Θ=θi=10000−2(i−1)/d,i∈[1,2,…,d/2] \Theta = {\theta_i = 10000^{-2(i-1)/d}, \quad i \in [1, 2, \dots, d/2]}
高性能向量化 PyTorch 实现
实际中为每个头和每个 token 显式构建完整块对角旋转矩阵会很慢。我们可以使用复数表示或实数切片技巧高效实现 RoPE:
RΘ,mx=x⊙cos(mΘ)+x~⊙sin(mΘ) R_{\Theta, m} \mathbf{x} = \mathbf{x} \odot \cos(m\Theta) + \tilde{\mathbf{x}} \odot \sin(m\Theta)
其中 x~=[−x2,x1,−x4,x3,… ]\tilde{\mathbf{x}} = [-x_2, x_1, -x_4, x_3, \dots] 。
以下是简洁、生产级的 PyTorch 实现:
python
import torch
import torch.nn as nn
class RotaryPositionalEmbedding(nn.Module):
"""
Rotary Position Embedding (RoPE) for multi-head attention.
Applies 2D plane rotations across d_model / 2 feature pairs.
"""
def __init__(self, dim: int, max_seq_len: int = 8192, base: float = 10000.0):
super().__init__()
self.dim = dim
self.base = base
# Compute frequency bands theta_i = 10000^(-2(i-1)/d)
inv_freq = 1.0 / (self.base ** (torch.arange(0, dim, 2).float() / dim))
self.register_buffer("inv_freq", inv_freq, persistent=False)
# Precompute cosine and sine cache for max_seq_len
self._build_cache(max_seq_len)
def _build_cache(self, max_seq_len: int):
t = torch.arange(max_seq_len, dtype=self.inv_freq.dtype)
# Outer product: [seq_len, dim / 2]
freqs = torch.outer(t, self.inv_freq)
# Duplicate frequencies to match feature dimensions: [seq_len, dim]
emb = torch.cat((freqs, freqs), dim=-1)
self.register_buffer("cos_cached", emb.cos(), persistent=False)
self.register_buffer("sin_cached", emb.sin(), persistent=False)
def _rotate_half(self, x: torch.Tensor) -> torch.Tensor:
"""Splits vector in half and negates/swaps halves: [-x2, x1]"""
x1 = x[..., : self.dim // 2]
x2 = x[..., self.dim // 2 :]
return torch.cat((-x2, x1), dim=-1)
def forward(self, x: torch.Tensor, seq_len: int) -> torch.Tensor:
"""
Args:
x: Tensor of shape [batch_size, num_heads, seq_len, head_dim]
seq_len: Current sequence length
Returns:
Tensor with RoPE applied to Query or Key
"""
cos = self.cos_cached[:seq_len, :].unsqueeze(0).unsqueeze(1) # [1, 1, seq_len, dim]
sin = self.sin_cached[:seq_len, :].unsqueeze(0).unsqueeze(1) # [1, 1, seq_len, dim]
# Apply elementwise rotation formula
return (x * cos) + (self._rotate_half(x) * sin)
Enter fullscreen mode Exit fullscreen mode
0 Comments
Log in to join the conversation.No comments yet. Be the first to share your thoughts.