Learn

旋转位置编码 RoPE

S=QKTdkS = \frac{QK^{\mathsf{T}}}{\sqrt{d_k}}

位置信息编码在 embedding 里:X=E+PX = E + P。记录了绝对信息

绝对位置编码的不足:

  1. 训练和推理的长度不一样,token的绝对位置也不一样,外推性差

相对位置编码RoPE不修改embedding,而是通过旋转KQ:

S=Q~K~TdkS = \frac{\tilde{Q}\tilde{K}^{\mathsf{T}}}{\sqrt{d_k}}

1. 计算RoPE

假设输入(不带绝对位置信息):

X=E∈R3×4X = E \in \mathbb{R}^{3 \times 4}

计算kqv:

Q=XWQ,K=XWK,V=XWVQ = XW_Q,\quad K = XW_K,\quad V = XW_V Q,K,V∈R3×2Q, K, V \in \mathbb{R}^{3 \times 2}

得到:

Q=[q0q1q2],K=[k0k1k2]Q = \begin{bmatrix} q_0 \\ q_1 \\ q_2 \end{bmatrix}, \quad K = \begin{bmatrix} k_0 \\ k_1 \\ k_2 \end{bmatrix}

每个向量都是平面上的一个点。位置 tt 对应角度 θt=t⋅ω\theta_t = t\cdot\omega,用旋转矩阵把它转过去:

Rt=[cos⁡θt−sin⁡θtsin⁡θtcos⁡θt]R_t = \begin{bmatrix} \cos\theta_t & -\sin\theta_t \\ \sin\theta_t & \cos\theta_t \end{bmatrix} q~t=Rt qt,k~t=Rt kt\tilde{q}_t = R_t\, q_t,\qquad \tilde{k}_t = R_t\, k_t

三个位置分别转:

q~0=R0q0,q~1=R1q1,q~2=R2q2\tilde{q}_0 = R_0 q_0,\quad \tilde{q}_1 = R_1 q_1,\quad \tilde{q}_2 = R_2 q_2

KK 同样操作。VV 不旋转。

然后照常打分:

S=Q~K~Tdk,A=softmax⁡(S+M),O=AVS = \frac{\tilde{Q}\tilde{K}^{\mathsf{T}}}{\sqrt{d_k}}, \quad A = \operatorname{softmax}(S + M), \quad O = AV

多头时每个头各自转自己的 QiQ_i、KiK_i,再代入第 2 章的 head⁡i\operatorname{head}_i。

qmq_m 转了 mωm\omega,knk_n 转了 nωn\omega。两点积里两个旋转合成一次,转角只剩 (n−m)ω(n-m)\omega。

2. 头宽不是 2 的时候

实际 dkd_k 常常是 64、128。不能整段当作一个平面,就把最后一维拆成 dk/2d_k/2 对,每一对按上面的二维方式转,各用一个频率:

ωi=base−2i/dk\omega_i = \mathrm{base}^{-2i/d_k}

高频对近邻敏感,低频对更远的间隔仍有区分度。计算步骤不变:先投影,再按对旋转 QQ、KK,再打分。

3. Prefill 与 Decode

公式与第 1 章相同,只是 QQ、KK 先转再进 QKTQK^{\mathsf{T}}。

Prefill 三个 token 用 R0,R1,R2R_0,R_1,R_2,一次转完再算整张 SS。

Decode 第 tt 步只来一个新 token,旋转用全局下标 tt。prompt 长度为 3 时,第一个生成 token 用 R3R_3,不是 R0R_0。写入 KV Cache 的 kk 已经转过,不必重算;当前 q~t\tilde{q}_t 直接和缓存里的 K~≤t\tilde{K}_{\le t} 打分。

4. 小结

  • 位置从「加在 XX 上」改成「旋转 QQ、KK」。
  • VV 和 mask 不变。
  • 点积依赖相对间隔 m−nm-n。
  • Decode 按全局位置旋转;缓存中的 KK 已是旋转后的结果。

下一章把多头末尾提到的 GQA 展开:Query 头多于 KV 组时,公式与 KV Cache 形状怎么变。