RoPE详解:从二维旋转到PyTorch手写旋转位置编码
前面的 Transformer 和 Mini-GPT 文章分别使用了固定正弦位置编码与可学习绝对位置向量。它们都在解决同一个问题:
Self-Attention 可以比较 token 内容,却不会天然知道 token 的先后顺序。
这篇文章继续研究另一种位置编码:RoPE(Rotary Position Embedding,旋转位置编码)。
RoPE 的做法很特别。它不把位置向量直接加到 token 表示上,而是在每个注意力头内部,把 Query 和 Key 按位置旋转不同角度:
token hidden state
-> 线性投影得到 Q、K、V
-> 根据 position 旋转 Q 和 K
-> 计算 QK^T
-> Softmax
-> 对没有旋转的 V 加权求和
它的核心价值可以浓缩为一句话:
用绝对位置决定旋转角度,让 Query 与 Key 的点积自然只显式依赖相对位置。
本文会从二维向量开始,逐步回答下面的问题:
- 为什么旋转能够表示位置?
- 为什么两个绝对位置最后会变成一个相对距离?
- 为什么只旋转 Q 和 K,不旋转 V?
- 高维向量怎样分组旋转?
base=10000与不同频率分别代表什么?- 怎样用 PyTorch 从零实现 RoPE?
- 怎样把 RoPE 接入 Mini-GPT、Causal Attention 与 KV Cache?
- RoPE 是否意味着模型可以无限外推到任意长度?
1. 先回顾:Attention 为什么需要位置信息
标准 Self-Attention 为:
$$
Q=XW_Q,\qquad K=XW_K,\qquad V=XW_V
$$
$$
\operatorname{Attention}(Q,K,V)
=\operatorname{softmax}\left(\frac{QK^T}{\sqrt{d_h}}+M\right)V
$$
其中:
- $X$ 是输入隐藏状态。
- $d_h$ 是单个注意力头的维度。
- $M$ 可以是 Causal Mask 或 Padding Mask。
假设序列中有三个 token:
我 喜欢 苹果
如果交换顺序:
苹果 喜欢 我
在不提供任何位置信息时,Self-Attention 仍然只看到三组内容向量。它能够计算“谁和谁更相关”,却缺少可靠的顺序坐标。
需要注意:
Causal Mask 不等于位置编码。
Causal Mask 只规定位置 $m$ 能不能看到位置 $n$:
$$
M_{m,n}=
\begin{cases}
0,&n\le m\
-\infty,&n>m
\end{cases}
$$
它能阻止模型偷看未来,却没有直接告诉模型两个可见 token 相隔 1 个位置还是 100 个位置。
2. 常见位置编码怎样注入位置
先把 RoPE 与前文出现过的方法放在一起比较。
2.1 固定正弦位置编码
原始 Transformer 构造一个与 token embedding 同维度的位置向量:
$$
z_m=x_m+p_m
$$
位置向量使用不同频率的正弦与余弦函数生成,并在进入第一层 Transformer 前与输入相加。
2.2 可学习绝对位置向量
前一篇 Mini-GPT 使用:
tok_emb = token_embedding(idx)
pos_emb = position_embedding(torch.arange(T))
x = tok_emb + pos_emb
每个位置对应一个可学习向量。优点是直观,缺点是位置表有固定大小,并且相对关系需要模型从数据中自己学出来。
2.3 RoPE
RoPE 不再执行:
token embedding + position embedding
而是在每一层 Attention 内执行:
Q_m = R_m(W_Q x_m)
K_m = R_m(W_K x_m)
V_m = W_V x_m
其中 $R_m$ 是由位置 $m$ 决定的旋转矩阵。
| 方法 | 位置注入在哪里 | 主要操作 | 是否直接影响 QK 匹配 |
|---|---|---|---|
| 固定正弦编码 | Transformer 输入前 | 与输入相加 | 间接影响 |
| 可学习绝对位置 | Transformer 输入前 | 查表后相加 | 间接影响 |
| RoPE | 每层 Attention 的 Q、K | 按位置旋转 | 直接影响 |
3. RoPE 的第一块积木:二维旋转
先暂时忘记高维张量,只看二维向量:
$$
x=
\begin{bmatrix}
x_1\
x_2
\end{bmatrix}
$$
将它逆时针旋转角度 $\phi$,使用旋转矩阵:
$$
R(\phi)=
\begin{bmatrix}
\cos\phi&-\sin\phi\
\sin\phi&\cos\phi
\end{bmatrix}
$$
旋转结果为:
$$
R(\phi)x=
\begin{bmatrix}
x_1\cos\phi-x_2\sin\phi\
x_1\sin\phi+x_2\cos\phi
\end{bmatrix}
$$
3.1 一个最简单的数值例子
令:
$$
x=
\begin{bmatrix}
1\0
\end{bmatrix}
$$
旋转 $90^\circ$:
$$
R(90^\circ)=
\begin{bmatrix}
0&-1\
1&0
\end{bmatrix}
$$
所以:
$$
R(90^\circ)x=
\begin{bmatrix}
0\1
\end{bmatrix}
$$
继续旋转到 $180^\circ$:
$$
R(180^\circ)x=
\begin{bmatrix}
-1\0
\end{bmatrix}
$$
如果规定“每前进一个 token 就多旋转 $90^\circ$”,那么同一个内容向量在不同位置会变成:
| 位置 $m$ | 角度 $m\theta$ | 旋转后向量 |
|---|---|---|
| 0 | $0^\circ$ | $[1,0]$ |
| 1 | $90^\circ$ | $[0,1]$ |
| 2 | $180^\circ$ | $[-1,0]$ |
| 3 | $270^\circ$ | $[0,-1]$ |
这已经包含了 RoPE 最基础的直觉:
位置不是被加进向量,而是被编码成向量的相位。
3.2 旋转不会改变向量长度
旋转矩阵是正交矩阵:
$$
R(\phi)^TR(\phi)=I
$$
因此:
$$
|R(\phi)x|_2=|x|_2
$$
也就是说,RoPE 改变方向,但不单独放大或缩小 Q、K 的模长。内容强度仍然保留,位置主要通过方向关系参与点积。
4. 从绝对位置到相对位置:最关键的推导
假设位置 $m$ 的原始 Query 为 $q_m$,位置 $n$ 的原始 Key 为 $k_n$。
RoPE 根据各自绝对位置进行旋转:
$$
\widetilde q_m=R(m\theta)q_m
$$
$$
\widetilde k_n=R(n\theta)k_n
$$
Attention 分数中的点积为:
$$
\widetilde q_m^T\widetilde k_n
=\left(R(m\theta)q_m\right)^T
\left(R(n\theta)k_n\right)
$$
把转置展开:
$$
=q_m^TR(m\theta)^TR(n\theta)k_n
$$
旋转矩阵满足:
$$
R(\alpha)^T=R(-\alpha)
$$
以及:
$$
R(\alpha)R(\beta)=R(\alpha+\beta)
$$
所以:
$$
R(m\theta)^TR(n\theta)
=R(-m\theta)R(n\theta)
=R((n-m)\theta)
$$
最终得到:
$$
\boxed{
\widetilde q_m^T\widetilde k_n
=q_m^TR((n-m)\theta)k_n
}
$$
右侧的位置项只包含 $n-m$,而不再分别依赖 $m$ 和 $n$。
这就是 RoPE 最重要的数学性质:
Q 和 K 各自使用绝对位置旋转,但二者点积中的位置关系表现为相对距离。
4.1 符号方向为什么有时写成 $m-n$
有些资料会写成 $R((m-n)\theta)$,有些写成 $R((n-m)\theta)$。差异通常来自:
- 将 Query 放在左边还是右边。
- 点积推导采用行向量还是列向量。
- 旋转方向的定义不同。
真正不变的是:
分数依赖两个位置之差,而不是两个独立的绝对位置。
实现时只要 Q、K 使用同一套约定即可。
5. 手算例子一:验证相对距离性质
为了便于计算,令每个位置增加的角度为:
$$
\theta=90^\circ
$$
原始向量为:
$$
q=
\begin{bmatrix}
1\2
\end{bmatrix},\qquad
k=
\begin{bmatrix}
3\4
\end{bmatrix}
$$
现在让 Query 位于 $m=1$,Key 位于 $n=2$。
5.1 分别按绝对位置旋转
Query 旋转 $90^\circ$:
$$
\widetilde q_1=R(90^\circ)q
=
\begin{bmatrix}
0&-1\1&0
\end{bmatrix}
\begin{bmatrix}
1\2
\end{bmatrix}
=
\begin{bmatrix}
-2\1
\end{bmatrix}
$$
Key 旋转 $180^\circ$:
$$
\widetilde k_2=R(180^\circ)k
=
\begin{bmatrix}
-1&0\0&-1
\end{bmatrix}
\begin{bmatrix}
3\4
\end{bmatrix}
=
\begin{bmatrix}
-3\-4
\end{bmatrix}
$$
二者点积:
$$
\widetilde q_1^T\widetilde k_2
=(-2)\times(-3)+1\times(-4)
=2
$$
5.2 只用相对距离计算
两个位置的差为:
$$
n-m=2-1=1
$$
先把 Key 只旋转一个相对角度:
$$
R(90^\circ)k
=
\begin{bmatrix}
-4\3
\end{bmatrix}
$$
再与原始 Query 点积:
$$
q^TR(90^\circ)k
=1\times(-4)+2\times3
=2
$$
两种算法结果完全一致:
分别使用绝对位置旋转后点积 = 2
只使用相对距离旋转后点积 = 2
如果把位置整体向后平移 10:
m = 11
n = 12
n - m 仍然等于 1
位置项仍然相同。这种平移一致性正是相对位置表示需要的性质。
6. 手算例子二:从旋转一直算到 Attention 输出
现在模拟一个只有两个 token、单头维度为 2 的 Causal Attention。
设两个位置投影后的 Q、K 内容恰好相同:
$$
q_0=q_1=k_0=k_1=[1,0]
$$
Value 为:
$$
v_0=[10,0],\qquad v_1=[0,20]
$$
仍令每个位置旋转 $90^\circ$。
6.1 对 Q 和 K 应用 RoPE
位置 0 不旋转:
$$
\widetilde q_0=\widetilde k_0=[1,0]
$$
位置 1 旋转 $90^\circ$:
$$
\widetilde q_1=\widetilde k_1=[0,1]
$$
6.2 计算位置 1 的注意力分数
位置 1 可以看到位置 0 和位置 1。
对位置 0:
$$
s_{1,0}
=\frac{\widetilde q_1\cdot\widetilde k_0}{\sqrt2}
=\frac{[0,1]\cdot[1,0]}{\sqrt2}
=0
$$
对位置 1:
$$
s_{1,1}
=\frac{\widetilde q_1\cdot\widetilde k_1}{\sqrt2}
=\frac{[0,1]\cdot[0,1]}{\sqrt2}
=\frac1{\sqrt2}
\approx0.7071
$$
6.3 Softmax 得到权重
$$
\operatorname{softmax}([0,0.7071])
\approx[0.3302,0.6698]
$$
6.4 对未旋转的 Value 加权求和
$$
o_1=0.3302v_0+0.6698v_1
$$
$$
=0.3302[10,0]+0.6698[0,20]
$$
$$
\approx[3.3024,13.3952]
$$
完整过程为:
Q、K 内容相同
-> 因位置不同而旋转到不同方向
-> 点积包含相对位置信息
-> Softmax 得到位置相关权重
-> 权重对原始 V 聚合
这个例子也提前回答了一个常见问题:V 不必旋转,因为位置已经参与了“应该读取哪个 V、读取多少”的权重计算。
7. 高维向量怎样旋转
真实注意力头的维度不会只有 2。假设:
$$
d_h=8
$$
RoPE 把最后一维两两分组:
(x0, x1) 第 0 组
(x2, x3) 第 1 组
(x4, x5) 第 2 组
(x6, x7) 第 3 组
每一组都看成一个二维平面,并使用自己的频率旋转:
$$
(x_{2j},x_{2j+1})
\xrightarrow{\ m\theta_j\ }
(x’{2j},x’{2j+1})
$$
第 $j$ 组的计算为:
$$
x’{2j}
=x{2j}\cos(m\theta_j)
-x_{2j+1}\sin(m\theta_j)
$$
$$
x’{2j+1}
=x{2j}\sin(m\theta_j)
+x_{2j+1}\cos(m\theta_j)
$$
7.1 为什么 head dimension 必须是偶数
每次旋转需要两个坐标组成一个平面,因此完整 RoPE 通常要求:
$$
d_h\bmod2=0
$$
如果只旋转前 rotary_dim 个维度,那么需要 rotary_dim 为偶数;剩余维度可以保持不变,这称为 Partial RoPE。
8. 不同维度为什么使用不同频率
RoPE 常用频率定义为:
$$
\theta_j=\text{base}^{-\frac{2j}{d_h}},
\qquad
j=0,1,\ldots,\frac{d_h}{2}-1
$$
经典设置为:
$$
\text{base}=10000
$$
位置 $m$ 对第 $j$ 组产生的实际角度为:
$$
\phi_{m,j}=m\theta_j
$$
例如 $d_h=4$、base=10000:
$$
\theta_0=10000^0=1
$$
$$
\theta_1=10000^{-1/2}=0.01
$$
因此两个二维平面的角度分别是:
第 0 组:m × 1 rad
第 1 组:m × 0.01 rad
第 0 组变化快,更容易分辨邻近位置;第 1 组变化慢,可以在更长距离上保持平缓变化。
这很像同时使用多个刻度不同的时钟:
- 高频维度像秒针,对局部位置变化敏感。
- 低频维度像时针,跨越较长距离后仍能提供缓慢变化的信号。
多个频率一起工作,使模型可以表示不同尺度的相对距离。
8.1 base 不是最大长度
base=10000 不表示“模型最多只能处理 10000 个 token”。它控制的是频率的分布。
一般来说,增大 base 会让低频部分旋转得更慢,但上下文能力还取决于:
- 预训练时见过的长度。
- 频率缩放方案。
- 模型是否针对新长度继续训练。
- Attention、数据与评估任务。
不能只改一个 base 就断言上下文窗口已经可靠扩展。
9. 完整旋转矩阵
对偶数维向量,位置 $m$ 的旋转矩阵可以写成分块对角矩阵:
$$
R_m=
\begin{bmatrix}
R(m\theta_0)&0&\cdots&0\
0&R(m\theta_1)&\cdots&0\
\vdots&\vdots&\ddots&\vdots\
0&0&\cdots&R(m\theta_{d_h/2-1})
\end{bmatrix}
$$
每个小块都是:
$$
R(m\theta_j)=
\begin{bmatrix}
\cos(m\theta_j)&-\sin(m\theta_j)\
\sin(m\theta_j)&\cos(m\theta_j)
\end{bmatrix}
$$
真实代码不会为每个 token 构造巨大的矩阵 $R_m$。那样既浪费内存,也会执行大量无意义的零元素乘法。
实现只需要缓存或即时计算:
cos(position × frequency)
sin(position × frequency)
然后使用逐元素乘法完成每个二维对的旋转。
10. 复数视角:旋转其实是乘相位
二维向量可以写成复数:
$$
z=x_1+ix_2
$$
欧拉公式为:
$$
e^{i\phi}=\cos\phi+i\sin\phi
$$
把复数乘以 $e^{i\phi}$:
$$
ze^{i\phi}
$$
它的模长不变,相位增加 $\phi$,几何意义正是二维旋转。
因此 RoPE 也可以写成:
$$
\widetilde q_{m,j}=q_{m,j}e^{im\theta_j}
$$
$$
\widetilde k_{n,j}=k_{n,j}e^{in\theta_j}
$$
这不是另一种算法,而是同一旋转的另一种表达方式。后面会分别给出实数版和复数版 PyTorch 实现。
11. RoPE 在 Attention 中的准确位置
正确顺序是:
x: [B, T, C]
-> Q/K/V Linear
-> 拆成多个头 [B, H, T, D]
-> 对 Q、K 应用 RoPE
-> QK^T / sqrt(D)
-> Causal Mask
-> Softmax
-> Attention Weight @ V
-> 合并多个头
这里:
- $B$:batch size。
- $T$:序列长度。
- $C$:模型维度。
- $H$:Query 头数。
- $D=C/H$:每个 Query 头的维度。
11.1 RoPE 不是旋转 token embedding
常见错误是直接写:
x = apply_rope(x)
q, k, v = qkv(x).chunk(3, dim=-1)
标准 RoPE 的目标是 Attention 内的 Query 和 Key。更符合定义的顺序是:
q, k, v = qkv(x).chunk(3, dim=-1)
q, k = apply_rope(q, k, cos, sin)
11.2 每层都要应用
每个 Transformer Block 都会产生自己的 Q 和 K,因此每一层 Attention 都需要对本层的 Q、K 应用 RoPE。
同一组位置频率可以在层之间共享,不必为每层学习不同的位置表。
12. 为什么只旋转 Q 和 K
Attention 可以拆成两个阶段:
- $QK^T$ 决定“应该关注谁”。
- Attention 权重乘 $V$ 决定“从对方取回什么内容”。
RoPE 主要希望让匹配分数包含相对位置,所以旋转 Q、K 已经足够:
$$
\operatorname{softmax}
\left(
\frac{(R_mq_m)^T(R_nk_n)}{\sqrt{d_h}}
\right)V
$$
Value 本身不参与 Query-Key 相似度计算,通常保持不旋转。
可以把它理解为:
Q、K:结合内容和距离,决定读取地址与权重
V:保存真正需要被聚合的内容
这不是说 V 完全没有上下文信息。前一层输出本身已经包含上下文,当前层的位置影响也会通过注意力权重进入加权结果。
13. PyTorch 实现一:预计算 cos 和 sin
先实现频率表:
import torch
def precompute_rope(head_dim, max_seq_len, base=10000.0, device=None):
if head_dim % 2 != 0:
raise ValueError("head_dim must be even for RoPE")
pair_index = torch.arange(
0, head_dim, 2, dtype=torch.float32, device=device
)
inv_freq = base ** (-pair_index / head_dim)
positions = torch.arange(
max_seq_len, dtype=torch.float32, device=device
)
angles = torch.outer(positions, inv_freq)
cos = angles.cos()
sin = angles.sin()
return cos, sin
形状变化为:
pair_index: [D/2]
inv_freq: [D/2]
positions: [T]
angles: [T, D/2]
cos, sin: [T, D/2]
以 D=8, T=128 为例:
cos.shape = [128, 4]
sin.shape = [128, 4]
13.1 为什么变量叫 inv_freq
代码常写成:
inv_freq = 1.0 / (
base ** (torch.arange(0, head_dim, 2) / head_dim)
)
它与:
base ** (-torch.arange(0, head_dim, 2) / head_dim)
完全等价。
这里的“inverse frequency”不是对最终角度求倒数,而是用负指数构造从快到慢的一组角频率。
14. PyTorch 实现二:旋转二维对
输入张量形状设为:
q, k: [B, H, T, D]
先实现 $[-x_2,x_1]$:
def rotate_pairs(x):
x_pairs = x.reshape(*x.shape[:-1], -1, 2)
x_even = x_pairs[..., 0]
x_odd = x_pairs[..., 1]
rotated = torch.stack((-x_odd, x_even), dim=-1)
return rotated.flatten(-2)
如果:
x = [x0, x1, x2, x3]
那么:
rotate_pairs(x) = [-x1, x0, -x3, x2]
再应用公式:
$$
R(\phi)x=x\cos\phi+\operatorname{rotate_pairs}(x)\sin\phi
$$
def apply_rope(q, k, cos, sin, position_ids):
# cos/sin: [max_seq_len, D/2]
# position_ids: [B, T]
cos = cos[position_ids].repeat_interleave(2, dim=-1)
sin = sin[position_ids].repeat_interleave(2, dim=-1)
# [B, T, D] -> [B, 1, T, D],广播到所有头
cos = cos.unsqueeze(1)
sin = sin.unsqueeze(1)
q_float = q.float()
k_float = k.float()
q_rotated = q_float * cos + rotate_pairs(q_float) * sin
k_rotated = k_float * cos + rotate_pairs(k_float) * sin
return q_rotated.to(q.dtype), k_rotated.to(k.dtype)
14.1 为什么用 repeat_interleave(2)
预计算时每个二维对只有一个角度:
[phi0, phi1, phi2, phi3]
一对中的两个坐标要共享同一个角度,所以扩展成:
[phi0, phi0, phi1, phi1, phi2, phi2, phi3, phi3]
对应的 cos 和 sin 才能与最后一维 [D] 逐元素运算。
14.2 为什么先转成 FP32
长位置与高频角度在低精度下更容易产生数值误差。教学实现中可以先把 Q、K 和角度计算提升到 FP32,旋转完成后再转回原始 dtype:
q_float = q.float()
...
q_rotated.to(q.dtype)
这不是唯一实现方式,但比直接在低精度里计算大量三角函数更稳妥。
15. PyTorch 实现三:封装成 nn.Module
实际模型经常动态传入 position_ids,可以写成:
from torch import nn
class RotaryEmbedding(nn.Module):
def __init__(self, head_dim, base=10000.0):
super().__init__()
if head_dim % 2 != 0:
raise ValueError("head_dim must be even for RoPE")
pair_index = torch.arange(0, head_dim, 2).float()
inv_freq = base ** (-pair_index / head_dim)
self.register_buffer(
"inv_freq", inv_freq, persistent=False
)
def forward(self, x, position_ids):
# position_ids: [B, T]
angles = (
position_ids.float().unsqueeze(-1)
* self.inv_freq.view(1, 1, -1)
)
angles = torch.repeat_interleave(angles, 2, dim=-1)
cos = angles.cos().unsqueeze(1).to(x.dtype)
sin = angles.sin().unsqueeze(1).to(x.dtype)
return cos, sin
这里不预先绑定最大序列长度,而是根据本批次的 position_ids 动态计算角度。
15.1 为什么使用 register_buffer
inv_freq:
- 不是通过梯度学习的参数。
- 需要随模型移动到 CPU 或 GPU。
- 不应该进入优化器。
所以它适合注册为 buffer,而不是 nn.Parameter。
persistent=False 表示不把它写入 state_dict,因为它可以由 head_dim 和 base 重新计算。如果希望完整保存,也可以省略该参数。
16. 复数版 PyTorch 实现
复数实现更贴近“乘相位”的数学表达:
def precompute_freqs_cis(head_dim, max_seq_len, base=10000.0):
pair_index = torch.arange(0, head_dim, 2).float()
inv_freq = base ** (-pair_index / head_dim)
positions = torch.arange(max_seq_len).float()
angles = torch.outer(positions, inv_freq)
return torch.polar(torch.ones_like(angles), angles)
def apply_rope_complex(x, freqs_cis):
# x: [B, T, H, D]
original_dtype = x.dtype
x_complex = torch.view_as_complex(
x.float().reshape(*x.shape[:-1], -1, 2)
)
# freqs_cis: [T, D/2] -> [1, T, 1, D/2]
freqs_cis = freqs_cis.view(
1, x.shape[1], 1, x.shape[-1] // 2
)
x_rotated = x_complex * freqs_cis
x_out = torch.view_as_real(x_rotated).flatten(-2)
return x_out.to(original_dtype)
对应关系是:
[x0, x1] -> x0 + i*x1
乘 exp(i*phi)
再拆回 [real, imag]
实数版和复数版结果应相同。工程中选择哪一种,取决于框架支持、内核、可读性与目标硬件。
17. 把 RoPE 接入多头因果注意力
下面给出一个完整、可运行的教学实现:
import torch
import torch.nn.functional as F
from torch import nn
def rotate_pairs(x):
x_pairs = x.float().reshape(*x.shape[:-1], -1, 2)
x_even = x_pairs[..., 0]
x_odd = x_pairs[..., 1]
return torch.stack((-x_odd, x_even), dim=-1).flatten(-2)
def apply_rotary(q, k, cos, sin):
q_float = q.float()
k_float = k.float()
q_out = q_float * cos + rotate_pairs(q_float) * sin
k_out = k_float * cos + rotate_pairs(k_float) * sin
return q_out.to(q.dtype), k_out.to(k.dtype)
class RotaryEmbedding(nn.Module):
def __init__(self, head_dim, base=10000.0):
super().__init__()
if head_dim % 2 != 0:
raise ValueError("head_dim must be even")
pair_index = torch.arange(0, head_dim, 2).float()
inv_freq = base ** (-pair_index / head_dim)
self.register_buffer(
"inv_freq", inv_freq, persistent=False
)
def forward(self, x, position_ids):
angles = (
position_ids.float().unsqueeze(-1)
* self.inv_freq.view(1, 1, -1)
)
angles = torch.repeat_interleave(angles, 2, dim=-1)
# [B, 1, T, D],头维自动广播
cos = angles.cos().unsqueeze(1)
sin = angles.sin().unsqueeze(1)
return cos, sin
class RoPECausalSelfAttention(nn.Module):
def __init__(
self,
n_embd,
n_head,
dropout=0.0,
rope_base=10000.0,
):
super().__init__()
if n_embd % n_head != 0:
raise ValueError("n_embd must be divisible by n_head")
self.n_head = n_head
self.head_dim = n_embd // n_head
if self.head_dim % 2 != 0:
raise ValueError("head_dim must be even for full RoPE")
self.qkv = nn.Linear(n_embd, 3 * n_embd, bias=False)
self.out_proj = nn.Linear(n_embd, n_embd, bias=False)
self.resid_dropout = nn.Dropout(dropout)
self.attn_dropout = dropout
self.rope = RotaryEmbedding(
head_dim=self.head_dim,
base=rope_base,
)
def forward(self, x, position_ids=None):
B, T, C = x.shape
if position_ids is None:
position_ids = torch.arange(
T, device=x.device
).unsqueeze(0).expand(B, -1)
q, k, v = self.qkv(x).chunk(3, dim=-1)
q = q.view(B, T, self.n_head, self.head_dim)
k = k.view(B, T, self.n_head, self.head_dim)
v = v.view(B, T, self.n_head, self.head_dim)
q = q.transpose(1, 2)
k = k.transpose(1, 2)
v = v.transpose(1, 2)
cos, sin = self.rope(q, position_ids)
q, k = apply_rotary(q, k, cos, sin)
y = F.scaled_dot_product_attention(
q,
k,
v,
is_causal=True,
dropout_p=(
self.attn_dropout if self.training else 0.0
),
)
y = y.transpose(1, 2).contiguous().view(B, T, C)
y = self.out_proj(y)
return self.resid_dropout(y)
17.1 与前一篇 Mini-GPT 相比改了什么
原来的输入:
x = token_embedding(idx) + position_embedding(pos)
改成 RoPE 后:
x = token_embedding(idx)
并在每层 Attention 中增加:
cos, sin = self.rope(q, position_ids)
q, k = apply_rotary(q, k, cos, sin)
也就是说:
- 删除输入处的绝对 Position Embedding。
- 保留 Token Embedding。
- 每层都在线性投影后旋转 Q、K。
- 后面的 Mask、Softmax、Value 聚合不变。
17.2 scaled_dot_product_attention 不会自动加入 RoPE
PyTorch 的 F.scaled_dot_product_attention 负责缩放点积、Mask、Softmax、Dropout 与 Value 聚合,但不会猜测模型采用哪一种位置编码。
所以必须在调用它之前完成:
q, k = apply_rotary(q, k, cos, sin)
18. 用断言验证实现是否正确
不要只看“代码能运行”。RoPE 至少应通过下面几类性质测试。
18.1 旋转保持范数
q_rot, k_rot = apply_rotary(q, k, cos, sin)
torch.testing.assert_close(
q.norm(dim=-1),
q_rot.norm(dim=-1),
rtol=1e-5,
atol=1e-5,
)
18.2 位置 0 不改变向量
因为:
$$
\cos0=1,\qquad\sin0=0
$$
所以位置 0 的旋转应为恒等变换:
torch.testing.assert_close(q_rot[:, :, 0], q[:, :, 0])
18.3 相同相对距离得到相同的位置作用
固定相同的基础 q、k,比较位置对:
(m, n) = (1, 3)
(m, n) = (5, 7)
它们都有:
$$
n-m=2
$$
因此旋转后的点积应接近相同。
score_13 = (q_at_1 * k_at_3).sum()
score_57 = (q_at_5 * k_at_7).sum()
torch.testing.assert_close(score_13, score_57)
18.4 实数版与复数版一致
使用同一输入和频率,两种实现输出应在浮点误差范围内一致。这是发现“维度配对方式不一致”最有效的测试之一。
19. RoPE 与 KV Cache
自回归生成时,每轮只新增一个 token。假设已经缓存位置 $0$ 到 $t-1$ 的 Key 和 Value:
cache_k: [B, H_kv, t, D]
cache_v: [B, H_kv, t, D]
生成位置 $t$ 时:
- 为新 token 计算 $q_t,k_t,v_t$。
- 使用绝对位置
position_id=t旋转 $q_t,k_t$。 - 把旋转后的 $k_t$ 放入 Key Cache。
- 把未旋转的 $v_t$ 放入 Value Cache。
- 用 $q_t$ 对全部缓存 Key 计算分数。
19.1 为什么缓存旋转后的 Key
历史 Key 的位置不会改变。位置 7 的 Key 在后续每一轮仍然属于位置 7,因此可以在第一次产生时就旋转并缓存,后续无需重复旋转。
伪代码如下:
q_new, k_new, v_new = project(x_new)
position_ids = torch.tensor([[cache_len]], device=x_new.device)
cos, sin = rope(q_new, position_ids)
q_new, k_new = apply_rotary(q_new, k_new, cos, sin)
cache_k = torch.cat([cache_k, k_new], dim=-2)
cache_v = torch.cat([cache_v, v_new], dim=-2)
output = attention(q_new, cache_k, cache_v)
19.2 不能每轮都把新 token 当作位置 0
有 KV Cache 时,本轮输入张量长度可能只有 1:
x_new.shape = [B, 1, C]
但这个 1 是本轮 token 数,不代表它的绝对位置为 0。
正确位置应该是:
position_id = cache_len
否则每个新 token 都会使用相同旋转角度,RoPE 的位置关系就被破坏。
20. Padding、左填充与 position_ids
不带 Padding 的 batch 可以直接使用:
position_ids = torch.arange(T, device=x.device)
position_ids = position_ids.unsqueeze(0).expand(B, -1)
但左填充时:
[PAD, PAD, 我, 爱, NLP]
[PAD, 今天, 天气, 很, 好]
简单的 [0,1,2,3,4] 会让第一条样本的真实 token 从位置 2 开始。很多实现会根据 Attention Mask 为每条样本构造连续位置:
position_ids = attention_mask.long().cumsum(dim=-1) - 1
position_ids.masked_fill_(attention_mask == 0, 0)
得到类似:
[0, 0, 0, 1, 2]
[0, 0, 1, 2, 3]
PAD 位置最终仍要被 Padding Mask 遮住。position_ids 与 Attention Mask 解决的是两个不同问题:
position_ids决定旋转角度。- Attention Mask 决定哪些 token 能参与注意力。
21. MHA、GQA 和 MQA 中怎样使用 RoPE
21.1 MHA
Multi-Head Attention 中:
Q: [B, H, T, D]
K: [B, H, T, D]
V: [B, H, T, D]
所有头共享相同位置频率表,但每个头拥有不同的 Q、K 内容,所以旋转后结果仍然不同。
21.2 GQA
Grouped-Query Attention 中 Query 头数多于 Key/Value 头数:
Q: [B, H_q, T, D]
K: [B, H_kv, T, D]
V: [B, H_kv, T, D]
只要 Q 和 K 的 head_dim 相同,同一组 [B,1,T,D] 的 cos、sin 就可以分别广播到 $H_q$ 和 $H_{kv}$。
21.3 MQA
Multi-Query Attention 只有一个 K/V 头,也仍然可以对:
K: [B, 1, T, D]
正常应用 RoPE。RoPE 与 MHA、GQA、MQA 描述的是不同维度:
- RoPE 负责位置关系。
- MHA/GQA/MQA 负责 Query 与 KV 头的组织方式。
22. Full RoPE 与 Partial RoPE
完整 RoPE 会旋转每个头的全部维度:
rotary_dim = head_dim
Partial RoPE 只旋转前一部分:
q_rot, q_pass = q[..., :rotary_dim], q[..., rotary_dim:]
k_rot, k_pass = k[..., :rotary_dim], k[..., rotary_dim:]
q_rot, k_rot = apply_rotary(q_rot, k_rot, cos, sin)
q = torch.cat([q_rot, q_pass], dim=-1)
k = torch.cat([k_rot, k_pass], dim=-1)
此时:
rotary_dim必须是偶数。- cos、sin 必须按照
rotary_dim生成。 - 剩余维度完全保留原值。
Full RoPE 与 Partial RoPE 都是合理设计,但必须与模型训练配置和已有权重保持一致。
23. 两种常见的维度配对约定
RoPE 代码里经常出现两种组织方法。
23.1 相邻配对
(x0, x1), (x2, x3), (x4, x5), ...
本文实数版和复数版采用这种方式。
23.2 前后半区配对
另一种实现会把向量拆成两半:
x1 = x[..., :D // 2]
x2 = x[..., D // 2:]
rotated = torch.cat((-x2, x1), dim=-1)
只要频率排列与权重训练方式配套,两种约定都能实现旋转思想。
但是:
不能只替换
rotate_half,却继续沿用不匹配的 cos、sin 排列或预训练权重。
从零训练时可以选择一种并保持一致;加载已有模型时必须严格遵循该模型的实现约定。
24. RoPE 与其他位置方法的比较
| 方法 | 注入方式 | 可训练位置参数 | 相对位置如何体现 | KV Cache 兼容性 |
|---|---|---|---|---|
| Learned Absolute | 输入加位置向量 | 有 | 模型间接学习 | 简单 |
| Sinusoidal | 输入加固定向量 | 无 | 可由三角关系间接推导 | 简单 |
| Relative Bias | Attention logits 加偏置 | 可有可无 | 直接按距离加分 | 通常良好 |
| ALiBi | logits 加距离线性偏置 | 通常无 | 直接惩罚距离 | 良好 |
| RoPE | 旋转 Q、K | 无 | QK 点积显式依赖相对相位 | 良好 |
RoPE 的优势包括:
- 不需要学习一张位置 embedding 表。
- 不改变 Q、K 的形状。
- 直接作用于 Attention 匹配分数。
- 与多头注意力和 KV Cache 容易结合。
- 可以通过频率缩放方法扩展上下文。
它也不是没有限制:
- 三角函数具有周期性。
- 超出训练长度后,模型未必理解陌生的相位分布。
- 不同实现的维度配对、缩放与位置约定可能不同。
- 长上下文能力不能只由“能算出 cos、sin”来判断。
25. RoPE 能无限外推吗
不能。
数学上,只要给出更大的位置 $m$,当然仍然能计算:
$$
\cos(m\theta_j),\qquad\sin(m\theta_j)
$$
但“程序可以计算”不等于“模型在该长度上训练过,也不等于模型能正确使用”。
当位置远超训练范围时,可能出现:
- 高频维度经历大量周期,产生训练中没见过的相位组合。
- 注意力分布发生异常变化。
- 短上下文能力与长上下文能力相互影响。
- 模型能够接收长输入,却无法准确检索或推理。
因此需要把三个概念分开:
张量长度支持:代码和显存能否容纳
位置编码支持:是否能为该位置计算信号
有效上下文能力:模型能否可靠利用远处信息
25.1 Position Interpolation
假设原训练长度为 $L$,目标长度为 $L’$,其中 $L’>L$。
位置插值的基本思想是把新位置压缩回原范围:
$$
m’=m\frac{L}{L’}
$$
然后使用 $m’$ 计算 RoPE 角度。这样不直接把模型推到远超训练范围的相位,但位置会变得更密集,通常仍需要继续训练或微调。
25.2 其他缩放方法
工程中还能看到 NTK-aware scaling、Dynamic NTK、YaRN 等方法。它们会以不同方式调整位置、频率或 Attention 缩放。
这些方法不是“随便改 base”的同义词。加载具体模型时应使用模型配置声明的 RoPE 类型和参数,不要凭名称猜测。
26. 计算量与参数量
RoPE 本身没有可训练参数。它增加的主要操作是:
- 生成或查询 cos、sin。
- 对 Q、K 做逐元素乘法和加法。
对于 Q、K 张量,额外计算量大致与:
$$
O(BHTD)
$$
同阶,而标准 Attention 分数矩阵的计算量为:
$$
O(BHT^2D)
$$
因此在长序列下,RoPE 通常不是 Attention 的主要复杂度来源。
如果预计算表,内存大致为:
$$
O(TD)
$$
如果动态计算,则可以减少固定缓存,但每次前向都要重新生成所需角度。实际系统会结合编译、缓存与上下文长度选择策略。
27. 常见错误
27.1 把 RoPE 加到 token embedding 上
标准用法是在线性投影之后旋转 Q 和 K,不是把旋转结果当作一张位置向量与输入相加。
27.2 旋转了 V,却忘记旋转 K
RoPE 的相对位置性质来自 Q、K 点积。只旋转 V 无法得到同样的推导。
27.3 在计算 QK 分数后才应用 RoPE
必须先旋转 Q、K,再做点积:
q, k = apply_rotary(q, k, cos, sin)
scores = q @ k.transpose(-2, -1)
27.4 head_dim 为奇数
完整两两配对会剩下一个坐标。应选择偶数 head_dim,或明确实现偶数大小的 Partial RoPE。
27.5 Softmax 维度错误
RoPE 不改变 Softmax 规则。Attention 仍然沿 Key 维归一化:
weights = torch.softmax(scores, dim=-1)
27.6 cos、sin 的形状无法正确广播
若 Q 为:
[B, H, T, D]
则 cos、sin 常整理成:
[B, 1, T, D]
不要把序列维与头维放反。
27.7 位置 0 出现变化
位置 0 应满足恒等旋转。如果不满足,通常说明:
- position offset 错误。
- sin/cos 排列错误。
- 二维配对方式错误。
27.8 KV Cache 每轮从位置 0 开始
新 token 的位置应为历史缓存长度,而不是本轮输入张量中的局部索引 0。
27.9 左填充后所有样本共用简单 arange
这可能让真实 token 的位置受 PAD 数量影响。应根据框架、训练方式和 Attention Mask 正确构造 position_ids。
27.10 混用相邻配对和前后半区配对
两种方式不能只替换一半。rotate 函数、cos/sin 排列与预训练权重必须来自同一约定。
27.11 把 base 当成最大上下文长度
base 控制频率分布,不直接等于最大 token 数。
27.12 只改缩放参数,不做长上下文评估
至少要检查:
- 原始短上下文困惑度是否退化。
- 长距离检索是否准确。
- 不同位置上的性能是否一致。
- 生成、代码与推理任务是否真正受益。
28. 高频问题
28.1 RoPE 是位置编码还是位置嵌入
中文资料中两种叫法都很常见。它没有一张通过梯度学习的位置向量表,因此从实现上更接近固定的位置编码机制;论文名称使用 Rotary Position Embedding。
28.2 RoPE 有可训练参数吗
标准 RoPE 没有。base、频率和位置规则属于超参数或固定计算。
28.3 每个注意力头的 RoPE 频率一样吗
通常相同。每个头共享频率规则,但 Q、K 内容不同,所以结果不会相同。
28.4 每一层的角度一样吗
同一位置、同一维度对应的基础角度通常一样,但每层 Q、K 投影结果不同,因此实际旋转后的向量不同。
28.5 还需要绝对 Position Embedding 吗
采用标准 RoPE 的 Decoder-only 模型通常不会再额外加入可学习绝对位置向量。但具体架构可以组合不同方法,最终应以模型定义为准。
28.6 RoPE 能替代 Causal Mask 吗
不能。RoPE 表示位置关系,Causal Mask 阻止访问未来,两者职责不同。
28.7 RoPE 会改变 Attention 的形状吗
不会:
旋转前 Q/K: [B, H, T, D]
旋转后 Q/K: [B, H, T, D]
28.8 为什么叫“旋转”而不是普通逐元素缩放
因为一对坐标同时使用 cos 和 sin 交叉混合:
$$
(x_1,x_2)
\rightarrow
(x_1\cos\phi-x_2\sin\phi,
x_1\sin\phi+x_2\cos\phi)
$$
这是真正的二维正交旋转,不是两个坐标各自乘一个系数。
29. 从 Prompt 到下一个 token:RoPE 参与了哪一步
承接前一篇 Mini-GPT,假设输入 Prompt:
春 天
Tokenizer 得到:
[id_春, id_天]
带 RoPE 的完整推理流程是:
1. Token Embedding
[B, T] -> [B, T, C]
2. 进入第 1 个 Transformer Block
LayerNorm -> Q/K/V 投影
3. 拆头
Q/K/V -> [B, H, T, D]
4. 生成 position_ids
[0, 1]
5. 应用 RoPE
位置 0 的 Q/K 不旋转
位置 1 的 Q/K 按多组频率旋转
6. Causal Attention
“天”的 Query 可以与“春”“天”的旋转后 Key 匹配
7. Attention 权重乘 V
聚合历史内容
8. MLP、残差与后续 Block
9. Final Norm + LM Head
得到最后位置的词表 logits
10. Softmax 与采样
例如预测“来”
RoPE 不会直接输出“来”。它的作用是改变每层 Attention 中的匹配分数,让模型在聚合“春”“天”等历史信息时能够利用相对位置。
当“来”被接回上下文:
春 天 来
下一轮的位置为 [0,1,2]。如果使用 KV Cache,只需要为位置 2 的新 token 计算并旋转新的 Q、K,再与历史缓存交互。
30. 学习与调试顺序
建议按下面顺序掌握 RoPE:
1. 手算二维旋转矩阵
2. 验证旋转前后范数不变
3. 推导 R_m^T R_n = R_(n-m)
4. 手算两位置的 QK 点积
5. 把 head_dim 拆成多个二维对
6. 生成多频率 cos/sin
7. 实现 rotate_pairs
8. 对 Q、K 应用 RoPE
9. 接入 Causal Self-Attention
10. 验证 position_ids 与 KV Cache offset
11. 再研究 Partial RoPE 和长上下文缩放
调试时优先打印:
print("q:", q.shape)
print("position_ids:", position_ids.shape)
print("cos:", cos.shape)
print("sin:", sin.shape)
print("q_rotated:", q_rotated.shape)
并检查:
print((q.norm(dim=-1) - q_rotated.norm(dim=-1)).abs().max())
该值应只包含很小的浮点误差。
31. 总结
RoPE 的完整主线可以压缩成:
hidden state
-> Q、K、V 投影
-> 按位置和多组频率生成角度
-> 把 Q、K 的最后一维两两配对
-> 每对执行二维旋转
-> 旋转后的 QK 点积只显式依赖相对位置差
-> Causal Mask + Softmax
-> 对未旋转的 V 加权求和
需要真正记住的核心包括:
- RoPE 旋转的是 Attention 中的 Q 和 K,不是把位置向量加到 token embedding。
- 二维旋转保持向量范数,并把位置编码成相位。
- $R_m^TR_n=R_{n-m}$ 让绝对位置旋转在点积中表现为相对距离。
- 高维向量被拆成多个二维对,每一对使用不同频率。
- 高频更敏感于局部变化,低频负责更长尺度的位置变化。
- V 通常不旋转,因为位置已经进入 QK 匹配权重。
- RoPE 不替代 Causal Mask,也不自动解决 Attention 的平方复杂度。
- KV Cache 中应缓存已经按绝对位置旋转的 Key,新 token 必须使用正确的位置偏移。
- 能计算超长位置的 cos、sin,不代表模型能可靠利用任意长上下文。
最后,用一条公式概括 RoPE:
$$
\boxed{
(R_mq_m)^T(R_nk_n)
=q_m^TR_{n-m}k_n
}
$$
左边使用两个 token 的绝对位置进行旋转,右边只剩二者的相对距离。这就是 RoPE 从“旋转”走向“相对位置”的关键桥梁。
完整示例代码:source/code/nlp-blog/rope.py
参考资料
- Jianlin Su et al., RoFormer: Enhanced Transformer with Rotary Position Embedding: https://arxiv.org/abs/2104.09864
- Ashish Vaswani et al., Attention Is All You Need: https://arxiv.org/abs/1706.03762
- Hugo Touvron et al., LLaMA: Open and Efficient Foundation Language Models: https://arxiv.org/abs/2302.13971
- Meta Llama,
apply_rotary_embreference implementation: https://github.com/meta-llama/llama/blob/main/llama/model.py - PyTorch Documentation,
scaled_dot_product_attention: https://docs.pytorch.org/docs/stable/generated/torch.nn.functional.scaled_dot_product_attention.html - Shouyuan Chen et al., Extending Context Window of Large Language Models via Positional Interpolation: https://arxiv.org/abs/2306.15595
- Bowen Peng et al., YaRN: Efficient Context Window Extension of Large Language Models: https://arxiv.org/abs/2309.00071
