跳转至

2.手撕旋转位置编码(RoPE)

推导 RoPE 的二维旋转矩阵

题目:写出在二维空间中,向量旋转角度 θ 的旋转矩阵形式,并说明 RoPE 中 θ 是如何与位置 m 关联的。

1.1 二维旋转矩阵

在二维平面上,将一个列向量 [x; y] 逆时针旋转角度 θ,其变换矩阵为:

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

作用于向量 v = [x, y]^T,得到 v' = R(θ) v = [x cos θ - y sin θ, x sin θ + y cos θ]^T。该操作保持了向量的模长,即 ||v'|| = ||v||

1.2 RoPE 中 θ 与位置 m 的关联

RoPE 的核心思想是将位置信息编码为旋转角度。对于序列中的第 m 个 token,其对应的查询向量 q_m 和键向量 k_n 在二维子空间中被旋转一个与位置成比例的角。

具体地,RoPE 将 d 维向量两两分组,每组对应一个不同的旋转频率。对于第 i 组(i = 0, 1, ..., d/2 - 1),其旋转角度为:

θ_i = 10000^{-2i/d}

位置 m 处的旋转角度为 m * θ_i。因此,位置 m 的向量在该子空间中的旋转矩阵为:

R(m, i) = [[cos(m θ_i), -sin(m θ_i)],
           [sin(m θ_i),  cos(m θ_i)]]

这样,q_mk_n 在第 i 个子空间的内积将包含相对位置信息 (m-n) θ_i(见后续等价性证明)。通过为不同维度组设置不同频率,模型可以同时捕捉短距离和长距离的依赖关系。


RoPE 的一般化形式推导

题目:将二维旋转推广到 d 维空间,写出 RoPE 对 query 和 key 向量施加位置编码的完整公式(分块对角矩阵形式)。

2.1 分块对角矩阵

对于 d 维向量(d 为偶数),RoPE 将其拆分为 d/2 个二维子空间,每个子空间应用独立的旋转。总体旋转矩阵为分块对角阵:

R(m) = diag(R(m,0), R(m,1), ..., R(m, d/2-1))

其中每个 2×2 块为:

R(m, i) = [[cos(m θ_i), -sin(m θ_i)],
           [sin(m θ_i),  cos(m θ_i)]]

这里 θ_i = 10000^{-2i/d}i = 0,1,...,d/2-1

2.2 对 Q 和 K 施加位置编码

对于位置 m 的查询向量 q_m ∈ R^d 和位置 n 的键向量 k_n ∈ R^d,RoPE 操作定义为:

q'_m = R(m) q_m
k'_n = R(n) k_n

具体计算时,通常采用逐元素乘加的方式(而非构造庞大的矩阵),将 q_m 的元素两两分组 (x_0, x_1), (x_2, x_3), ...,然后进行旋转变换:

x'_0 = x_0 cos(m θ_0) - x_1 sin(m θ_0)
x'_1 = x_0 sin(m θ_0) + x_1 cos(m θ_0)

依此类推。

2.3 完整公式表达

q_m = [q_0, q_1, ..., q_{d-1}],则旋转后:

q'_{2i} = q_{2i} cos(m θ_i) - q_{2i+1} sin(m θ_i)
q'_{2i+1} = q_{2i} sin(m θ_i) + q_{2i+1} cos(m θ_i)

k_n 同理。这个操作可以在实际代码中通过向量化的方式高效实现。


解释 RoPE 的远程衰减性质

题目:证明或直观解释:为什么 RoPE 可以使得内积随着相对距离的增大而衰减?给出公式推导。

3.1 衰减性的直观解释

RoPE 后的 query 和 key 内积为:

(q'_m)^T (k'_n) = (R(m) q_m)^T (R(n) k_n) = q_m^T R(m)^T R(n) k_n

由于旋转矩阵是正交矩阵,有 R(m)^T = R(-m),且 R(-m) R(n) = R(n - m),因此内积仅依赖于相对位置 Δ = n - m

(q'_m)^T (k'_n) = q_m^T R(Δ) k_n

展开后,对于第 i 个子空间,其贡献为:

q_{2i} k_{2i} cos(Δ θ_i) + q_{2i} k_{2i+1} sin(Δ θ_i) - q_{2i+1} k_{2i} sin(Δ θ_i) + q_{2i+1} k_{2i+1} cos(Δ θ_i)

对于随机的 q, k 分量,期望上该值随着 |Δ| 增大而振荡衰减,因为高频分量(i 大,θ_i 小)在长距离上迅速振荡,内积的平均值趋近于零;低频分量(i 小,θ_i 大)则提供主要的远程依赖。这种特性使得模型能自然学习到距离相关的权重衰减。

3.2 数学分析

可证明,当 qk 的每个分量独立同分布、均值为零时,E[(q'_m)^T (k'_n)] = 0,且方差随 |Δ| 增大而减小。具体来说,每个子空间的期望内积值与 cos(Δ θ_i) 成正比,随着 |Δ| 增大,多个不同频率的余弦项叠加会导致整体内积幅值衰减。这就是 RoPE 的远程衰减性质。


计算 RoPE 的基频选择

题目:给定维度 d=128,最大序列长度 L=2048,手动计算 RoPE 中频率 θ_i = 10000^{-2i/d} 的前 5 个频值,并解释基频 10000 的作用。

4.1 计算前 5 个频率

d=128,则 i = 0,1,2,3,4。 公式:θ_i = 10000^{-2i/d} = 10000^{-i/64}

  • i=0: θ_0 = 10000^{0} = 1.0

  • i=1: θ_1 = 10000^{-1/64} = e^{- (ln 10000)/64} ≈ e^{-9.2103/64} ≈ e^{-0.1439} ≈ 0.866

  • i=2: θ_2 = 10000^{-2/64} = 10000^{-1/32} ≈ e^{-0.2878} ≈ 0.75

  • i=3: θ_3 = 10000^{-3/64} ≈ e^{-0.4317} ≈ 0.649

  • i=4: θ_4 = 10000^{-4/64} = 10000^{-1/16} ≈ e^{-0.5756} ≈ 0.562

所以前五个频率约为 1.0, 0.866, 0.75, 0.649, 0.562。这些频率随着 i 增大而减小,即越靠后的维度组旋转速度越慢,捕捉更长的距离依赖。

4.2 基频 10000 的作用

基频 10000 决定了频率的整体尺度。较大的基频使得低频部分变化更缓慢,能支持更长的上下文窗口。RoPE 的外推能力依赖于基频的选择:增大基频(如 1000000)可以显著扩展有效上下文长度,而不需要微调。这是最近 NTK-aware 插值等方法的核心思想。


手写 RoPE 的 PyTorch 实现

题目:实现函数 apply_rotary_pos_emb(q, k, cos, sin, position_ids),要求支持 batching,且 cos/sin 为预先计算的固定频率表。

5.1 实现代码

import torch

def apply_rotary_pos_emb(q, k, cos, sin, position_ids):
    """
    q, k: (batch, num_heads, seq_len, head_dim)
    cos, sin: (max_seq_len, head_dim)  预计算的旋转因子
    position_ids: (batch, seq_len) 或 (seq_len,)
    """
    # 根据 position_ids 取出对应位置的 cos, sin
    cos = cos[position_ids].unsqueeze(0).unsqueeze(2)  # (1, batch, 1, seq_len, head_dim)
    sin = sin[position_ids].unsqueeze(0).unsqueeze(2)
    # 如果 q,k 是 (batch, heads, seq, dim),则 cos/sin 形状需适配
    cos = cos.unsqueeze(1)  # (1, 1, 1, seq_len, dim) -> 广播
    sin = sin.unsqueeze(1)

    # 将 q,k 拆分为两半
    q_rot = q.reshape(*q.shape[:-1], -1, 2)
    k_rot = k.reshape(*k.shape[:-1], -1, 2)
    # 应用旋转
    q_embed = torch.stack([
        q_rot[..., 0] * cos - q_rot[..., 1] * sin,
        q_rot[..., 0] * sin + q_rot[..., 1] * cos
    ], dim=-1).flatten(-2)
    k_embed = torch.stack([
        k_rot[..., 0] * cos - k_rot[..., 1] * sin,
        k_rot[..., 0] * sin + k_rot[..., 1] * cos
    ], dim=-1).flatten(-2)
    return q_embed, k_embed

更简洁的实现(利用复数)见后面第10题。这里的实现采用实数方式,将 head_dim 两两分组,直接相乘累加。

5.2 注意事项

  • cossin 应当预计算好全部位置的频率表,避免每次重复计算。

  • 支持 batch 和 multi-head 时,确保广播维度正确。常见输入形状是 (batch, num_heads, seq_len, head_dim),position_ids 形状 (batch, seq_len)

  • RoPE 通常只应用于 Q 和 K,不用于 V。


手写 RoPE 的预计算频率表生成

题目:编写函数 precompute_freqs_cis(dim, max_seq_len, theta=10000.0),输出复数形式的旋转因子 cis(即 cos + i*sin)。

6.1 实现代码

import torch

def precompute_freqs_cis(dim, max_seq_len, theta=10000.0):
    # 计算频率
    freqs = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim))  # (dim/2,)
    t = torch.arange(max_seq_len)  # (max_seq_len,)
    freqs = torch.outer(t, freqs)  # (max_seq_len, dim/2)
    # 转为复数 cis
    freqs_cis = torch.polar(torch.ones_like(freqs), freqs)  # (max_seq_len, dim/2)
    return freqs_cis

6.2 说明

  • freqs 的最后一维是 dim/2,对应每个二维子空间的频率。

  • 使用 torch.polar 直接创建复数:abs=1, angle=freqs,即 e^{i * freqs}

  • 得到的 freqs_cis 可以后续与 query/key 的复数形式直接相乘实现旋转。复数方式比实数分拆更简洁且高效(依赖底层优化)。


在多头自注意力中集成 RoPE

题目:修改标准注意力代码,在计算 Q、K 之后、计算注意力分数之前,对 Q、K 施加 RoPE,写出修改后的完整 attention_forward 函数。

7.1 实现代码

假设已有预计算好的 freqs_cis(形状 (max_seq_len, dim/2))和 position_ids

def attention_forward_rope(self, x, freqs_cis, position_ids, mask=None):
    B, seq_len, _ = x.shape
    # 投影 QKV
    qkv = self.qkv_proj(x)
    q, k, v = qkv.chunk(3, dim=-1)
    # 分头
    q = q.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
    k = k.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
    v = v.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2)

    # 应用 RoPE
    q, k = apply_rotary_pos_emb(q, k, freqs_cis, position_ids)  # 见第5题或复数版本

    # 缩放点积注意力
    scale = self.head_dim ** 0.5
    attn_scores = torch.matmul(q, k.transpose(-2, -1)) / scale
    if mask is not None:
        attn_scores = attn_scores.masked_fill(mask == 0, float('-inf'))
    attn_weights = F.softmax(attn_scores, dim=-1)
    attn_output = torch.matmul(attn_weights, v)

    # 合并多头
    attn_output = attn_output.transpose(1, 2).contiguous().view(B, seq_len, -1)
    return self.out_proj(attn_output)

7.2 关键点

  • RoPE 施加于分头之后,即形状 (B, num_heads, seq_len, head_dim)

  • apply_rotary_pos_emb 需要适配该形状(可内部处理)。

  • 在推理时,freqs_cis 可根据实际序列长度截取。


RoPE 与绝对位置编码(正弦/可学习)的对比

题目:写出两者在实现上的关键差异,并现场改写一个使用绝对位置编码的注意力模块,改为使用 RoPE。

8.1 关键差异

特性 绝对位置编码 (APE) RoPE
编码方式 将位置向量直接加到词嵌入上 通过旋转矩阵作用于 Q 和 K
依赖关系 内积依赖于绝对位置 内积仅依赖于相对位置
外推能力 差(固定最大长度) 较好(可通过插值扩展)
参数量 可学习参数矩阵 (max_len, dim) 无额外参数
实现复杂度 简单(相加) 较复杂(需旋转操作)

8.2 改写示例

原始 APE 注意力:

def forward_ape(self, x, pos_emb):
    x = x + pos_emb  # 加入位置信息
    # 然后进行注意力...

改为 RoPE:

  • 移除 pos_emb 加法,不再需要位置嵌入表。

  • 在 Q、K 投影后加入 RoPE 旋转步骤(如前题所示)。

  • 保留预计算的 freqs_cis


推导 RoPE 与相对位置编码的等价性

题目:证明施加 RoPE 后的 Q、K 内积只依赖于相对位置 m-n,写出推导步骤。

9.1 推导

q_m 为位置 m 的原始查询向量,k_n 为位置 n 的原始键向量。RoPE 后的内积:

<q'_m, k'_n> = (R(m) q_m)^T (R(n) k_n)
            = q_m^T R(m)^T R(n) k_n

由于旋转矩阵是正交矩阵,R(m)^T = R(-m)。又因为旋转具有可加性:R(a) R(b) = R(a+b),所以:

R(-m) R(n) = R(n - m)

因此:

<q'_m, k'_n> = q_m^T R(n - m) k_n

该结果仅与相对位置 Δ = n - m 有关,与绝对位置 m, n 无关。这正是相对位置编码的核心性质。

9.2 物理意义

注意力权重将反映 token 之间的相对距离,而非它们的绝对位置。这提升了模型对序列长度的泛化能力。


手写复数运算实现 RoPE

题目:利用复数乘法(torch.polar / torch.view_as_complex)实现 RoPE,比较与实数方式(cos/sin 直接乘加)的效率和代码简洁性。

10.1 复数实现

def apply_rotary_pos_emb_complex(q, k, freqs_cis, position_ids):
    """
    q, k: (batch, num_heads, seq_len, head_dim)
    freqs_cis: (max_seq_len, head_dim//2)  复数形式
    """
    # 根据 position_ids 索引频率
    cis = freqs_cis[position_ids]  # (batch, seq_len, head_dim//2)
    # 调整形状以匹配 q,k
    cis = cis.unsqueeze(1)  # (batch, 1, seq_len, head_dim//2)

    # 将 q,k 转为复数:每两个相邻维度作为实部/虚部
    q_ = torch.view_as_complex(q.float().reshape(*q.shape[:-1], -1, 2))
    k_ = torch.view_as_complex(k.float().reshape(*k.shape[:-1], -1, 2))

    # 复数乘法旋转
    q_rot = q_ * cis
    k_rot = k_ * cis

    # 转回实数
    q_embed = torch.view_as_real(q_rot).flatten(-2)
    k_embed = torch.view_as_real(k_rot).flatten(-2)
    return q_embed.type_as(q), k_embed.type_as(k)

10.2 对比

  • 代码简洁性:复数版本更短,语义更清晰,避免了手动的 cos/sin 拼接。

  • 效率:复数乘法底层调用优化过的复数运算库,在 GPU 上可能比手动拆分再乘加稍慢,但现代框架(如 PyTorch 2.0+)对其有良好支持。实际差异通常可忽略。

  • 可读性:复数形式更贴近数学定义,易于理解。

两种实现在效果上完全等价,选择哪种取决于个人偏好和具体硬件优化情况。在生产环境中,LLaMA 等模型早期采用实数版本,后续也有复数版本的应用。


处理 Q 和 K 维度不同的情况(如 GQA)如何实现 RoPE?广播逻辑。

在 GQA 或 MQA 中,Query 头数 h 大于 Key/Value 头数 g,但 Q 和 K 的 head_dim 通常保持一致。因此 RoPE 的旋转操作本身对每个头的 Q 和 K 独立进行,维度相同,所以可以直接分别旋转。问题在于 freqs_cis 的形状与广播:freqs_cis 的形状是 (max_seq_len, head_dim//2),与头数无关。

实现广播逻辑:

  • 分别对 Q 和 K 进行旋转,Q 形状为 (B, h, L, head_dim),K 形状为 (B, g, L, head_dim)

  • freqs_cis 扩展为 (1, 1, L, head_dim//2) 或通过 position_ids 索引得到 (B, 1, L, head_dim//2),然后分别应用到 Q 和 K 上。对于 Q 需要广播到 h 个头,K 广播到 g 个头。

  • 旋转操作本身与头数无关,只需要确保 freqs_cis 在头维度上广播(通过 unsqueeze)。

代码示例:

def apply_rotary_emb_gqa(q, k, freqs_cis, position_ids):
    # q: (B, num_q_heads, L, head_dim)
    # k: (B, num_kv_heads, L, head_dim)
    # freqs_cis: (max_seq_len, head_dim//2) 复数形式
    cis = freqs_cis[position_ids]  # (B, L, head_dim//2)
    # 为头维度增加维度并广播
    cis_q = cis.unsqueeze(1)  # (B, 1, L, head_dim//2) 自动广播到所有q头
    cis_k = cis.unsqueeze(1)  # 同样
    # 复数旋转
    q_ = torch.view_as_complex(q.float().reshape(*q.shape[:-1], -1, 2))
    k_ = torch.view_as_complex(k.float().reshape(*k.shape[:-1], -1, 2))
    q_rot = q_ * cis_q
    k_rot = k_ * cis_k
    q_out = torch.view_as_real(q_rot).flatten(-2).type_as(q)
    k_out = torch.view_as_real(k_rot).flatten(-2).type_as(k)
    return q_out, k_out

由于 head_dim 相同,旋转逻辑完全一致,因此 GQA/MQA 不影响 RoPE 的实现,只需在头维度上正确广播 freqs_cis 即可。


RoPE 在推理时如何与 KV Cache 结合?缓存的是旋转后的 K 还是旋转前的?

在自回归解码时,每步只生成一个新 token,需要计算该 token 的 Q、K、V,并将新 K、V 追加到 KV Cache 中。为了复用缓存的 K,RoPE 对 K 的旋转操作必须与位置相关。

方案选择:

  • 缓存旋转前的 K,每次使用时重新旋转:缓存的是未加位置编码的原始 K。在每一步,需要根据当前位置为所有缓存的 K 施加 RoPE(因为 K 的旋转角度依赖于其绝对位置)。但这会导致每一步都要对全部缓存进行旋转,计算量巨大,不可取。

  • 缓存旋转后的 K:将已经施加 RoPE 的 K(即 k_rot)直接存入缓存。这样在后续步中直接使用,无需重新旋转。这是标准做法。

具体实现:

在推理的每一步,对新 token 计算 Q 和 K,并根据当前 token 的绝对位置分别对 Q 和 K 施加 RoPE,然后将旋转后的 K 追加到 KV Cache。对于 Q,每次只计算当前 token 的 Q,只需旋转当前一个位置。因此 KV Cache 中存储的是已经旋转后的 K 和 V(V 无需旋转)。这样,注意力计算时,Q 和缓存的 K 都是旋转后的,内积自然包含了相对位置信息。

伪代码:

# step: 当前步数(从0开始)
position = step  # 当前token的绝对位置
q = self.q_proj(x)  # 新 token 的 Q
k = self.k_proj(x)
v = self.v_proj(x)
# 施加 RoPE
q_rot, k_rot = apply_rotary_pos_emb(q, k, freqs_cis, position)
# 更新 cache
self.k_cache[:, :, step:step+1] = k_rot
self.v_cache[:, :, step:step+1] = v
# 取出全部缓存 K, V(形状:B, H, step+1, D)
k_all = self.k_cache[:, :, :step+1]
v_all = self.v_cache[:, :, :step+1]
# 注意力计算
attn_output = scaled_dot_product_attention(q_rot, k_all, v_all)

缓存的是旋转后的 K。


从零实现 LLaMA 风格的 RoPE:LlamaRotaryEmbedding 类

参照 LLaMA 源码,其 RoPE 实现通过预计算复数形式的 freqs_cis,并提供了 forward 方法根据序列长度返回 cossin(或直接返回复数)。下面给出完整实现:

import torch
import torch.nn as nn

class LlamaRotaryEmbedding(nn.Module):
    def __init__(self, dim, max_position_embeddings=2048, base=10000.0, device=None):
        super().__init__()
        self.dim = dim          # head_dim
        self.max_position_embeddings = max_position_embeddings
        self.base = base
        # 计算 inv_freq: (dim/2,)
        inv_freq = 1.0 / (self.base ** (torch.arange(0, dim, 2).float().to(device) / dim))
        self.register_buffer("inv_freq", inv_freq, persistent=False)

    @torch.no_grad()
    def forward(self, x, position_ids):
        """
        x: 输入,仅用于获取 dtype 和 device
        position_ids: (batch_size, seq_len) 或 (seq_len,)
        返回 cos, sin,形状 (batch_size, seq_len, dim)
        """
        # 扩展维度以进行外积
        inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1)
        position_ids_expanded = position_ids[:, None, :].float()
        # 计算频率矩阵: (batch, dim/2, seq_len)
        freqs = (inv_freq_expanded @ position_ids_expanded).transpose(1, 2)  # (batch, seq_len, dim/2)
        # 拼接成 dim 维度(每个频率重复两次)
        emb = torch.cat((freqs, freqs), dim=-1)  # (batch, seq_len, dim)
        cos = emb.cos().to(x.dtype)
        sin = emb.sin().to(x.dtype)
        return cos, sin

使用时:

cos, sin = rotary_emb(x, position_ids)
# 在注意力中应用旋转
q_embed = (q * cos.unsqueeze(1)) + (rotate_half(q) * sin.unsqueeze(1))
k_embed = (k * cos.unsqueeze(1)) + (rotate_half(k) * sin.unsqueeze(1))

其中 rotate_half 将后半部分取负并交换前后半部分,实现二维旋转。


分析 RoPE 对长文本外推的限制,并介绍 NTK-aware 插值

外推失败的原因: RoPE 的频率范围从 1(i=0)到接近 0(i=d/2-1)。当序列长度超过训练时的最大长度,模型会遇到从未见过的 m * θ_i 值,尤其是高频分量(i 大,θ_i 小)在长距离上的旋转角度超出了训练分布,导致注意力混乱。

具体来说,θ_i = base^{-2i/d},最大频率为 1(对应周期为 2π ≈ 6.28),最小频率为 base^{-(d-2)/d} 接近 1/base。训练长度为 2048 时,模型学到的是相对距离在 [0, 2048] 范围内的模式。当距离扩大到 4096,对于较小的 θ_i(低频),Δ * θ_i 仍在训练范围内,但高频分量的 Δ * θ_i 会超出,使得这些分量的内积行为不可预测,破坏注意力分布。这导致困惑度急剧上升。

NTK-aware 插值: 核心思想是缩放频率,将高频分量的旋转速度降低,使得新长度下的最大相对距离对应的旋转角度仍落在训练分布内。具体做法是将 base 从 10000 增大到 10000 * α^{d/(d-2)},其中 α 是长度扩展倍数。这等价于将频率 θ_i 按照其索引 i 进行非线性缩放:高频(i 大)缩放更多,低频几乎不变。

公式:新频率 θ'_i = θ_i * (base_new / base)^{-2i/d},可以仅通过修改 base 实现。

实现:

def ntk_aware_rope_scale(inv_freq, scale_factor):
    # inv_freq: (dim/2,)
    # 采用 NTK-aware 缩放
    base = 10000.0
    new_base = base * scale_factor ** (dim / (dim - 2))
    inv_freq = 1.0 / (new_base ** (torch.arange(0, dim, 2).float() / dim))
    return inv_freq

手写 NTK-aware RoPE 插值

根据上题,实现一个函数,输入原始 inv_freq 和扩展倍数 scale,输出缩放后的 inv_freq

def ntk_aware_interpolate(inv_freq, scale, dim):
    """
    inv_freq: (dim/2,) 原始频率的倒数
    scale: 扩展倍数,例如 2 表示从 2048 扩展到 4096
    dim: head_dim
    """
    base = 10000.0
    # 新的 base 计算
    new_base = base * (scale ** (dim / (dim - 2)))
    # 重新生成 inv_freq
    new_inv_freq = 1.0 / (new_base ** (torch.arange(0, dim, 2).float() / dim))
    return new_inv_freq

在实际使用中,可以用新的 inv_freq 替换 RoPE 的 inv_freq,然后预计算 freqs_cis,无需重新训练模型。


RoPE 与 ALiBi 的异同

ALiBi 实现:在注意力分数上加一个与距离成线性比例的负偏置,斜率对每个头不同,无需位置编码。代码片段:

def get_alibi_slopes(num_heads):
    # 生成头相关的斜率
    return torch.tensor([2 ** (-8 / num_heads * i) for i in range(1, num_heads+1)])

# 在注意力分数上添加
alibi_bias = -slopes.view(-1,1,1) * distances.abs().unsqueeze(0)  # 因果单向时取负距离
attn_scores = attn_scores + alibi_bias

对比:

特性 RoPE ALiBi
施加位置 在 Q 和 K 上直接旋转 在注意力分数上加偏置
依赖关系 内积仅依赖于相对位置 仅依赖于相对距离(绝对值)
外推能力 较好,可通过插值扩展 天生支持外推,无长度限制
计算开销 需在 Q/K 上执行旋转操作 仅需加法,极其轻量
参数量 0 0(斜率由公式确定)
模型质量 在长序列上通常更好 在短序列上效果相当,长序列略逊

RoPE 对信息进行了旋转变换,保留了更多方向性信息;ALiBi 是纯粹的偏置,结构更简单。两者均可用于长序列,但 RoPE 更常用在现代大模型中,ALiBi 则在一些高效模型中出现。


在 RoPE 中混合可学习参数

可以引入可学习的缩放因子 γ_i 和偏移 β_i 来调整每个频率分量的旋转角度,公式变为:

θ'_i = γ_i * θ_i + β_i

或直接学习一个频率调整量。实现时,将 γ_iβ_i 设为 nn.Parameter,初始化 γ_i=1, β_i=0

代码修改:

class LearnableRoPE(LlamaRotaryEmbedding):
    def __init__(self, dim, max_position_embeddings=2048, base=10000.0):
        super().__init__(dim, max_position_embeddings, base)
        self.gamma = nn.Parameter(torch.ones(dim // 2))
        self.beta = nn.Parameter(torch.zeros(dim // 2))

    def forward(self, x, position_ids):
        inv_freq = self.inv_freq * self.gamma  # 缩放频率
        # 加上偏移会影响相位,但这里 inv_freq 调整后需要重新生成 cos/sin
        # 具体实现中可重新计算 freqs
        ...

可能的问题:

  • 学习参数可能破坏 RoPE 的远程衰减性质,导致训练不稳定。

  • 如果不加正则化,模型可能学出极端的频率,影响泛化。

  • 在微调阶段使用少量数据学习可能有效,但从头训练可能增加难度。 因此通常不引入可学习参数,除非在特定微调场景下并配合强正则化。


验证 RoPE 的对称性和周期性

测试脚本:

def test_rope_relative_property():
    dim = 64
    max_len = 100
    base = 10000.0
    freqs_cis = precompute_freqs_cis(dim, max_len, base)  # (max_len, dim//2)

    # 随机生成两个 query 和 key 向量
    B, H, L = 2, 4, 10
    q = torch.randn(B, H, L, dim)
    k = torch.randn(B, H, L, dim)
    # 随机选择一对位置 m, n
    m, n = 3, 7
    delta = n - m
    # 施加 RoPE
    def rotate(t, pos):
        cis = freqs_cis[pos].unsqueeze(0).unsqueeze(0)  # (1,1,dim//2)
        t_ = torch.view_as_complex(t.float().reshape(*t.shape[:-1], -1, 2))
        t_rot = t_ * cis
        return torch.view_as_real(t_rot).flatten(-2).type_as(t)

    qm_rot = rotate(q[:, :, m:m+1], m)
    kn_rot = rotate(k[:, :, n:n+1], n)
    # 内积
    score_mn = torch.einsum('b h d, b h d -> b h', qm_rot.squeeze(2), kn_rot.squeeze(2))

    # 改变绝对位置但保持相对距离:m'=2, n'=6 (Δ=4)
    m2, n2 = 2, 6
    qm2_rot = rotate(q[:, :, m2:m2+1], m2)
    kn2_rot = rotate(k[:, :, n2:n2+1], n2)
    score_m2n2 = torch.einsum('b h d, b h d -> b h', qm2_rot.squeeze(2), kn2_rot.squeeze(2))

    # 检查是否接近
    assert torch.allclose(score_mn, score_m2n2, atol=1e-4), "RoPE 内积不是仅依赖于相对位置!"
    print("测试通过:内积仅依赖于相对位置。")

    # 测试远程衰减:计算不同 Δ 的内积并观察趋势
    for delta in [1, 2, 4, 8, 16, 32, 64, 128]:
        # 随机生成 q 和 k,计算平均内积
        pass  # 可扩展

此脚本验证了 RoPE 的核心性质。


使用 Triton 或 CUDA 融合 RoPE 与注意力计算

融合的基本想法:在注意力 kernel 中,加载 Q 和 K 的块后,直接在片上计算 RoPE 的旋转,避免将旋转后的 Q、K 写回 HBM。伪代码(CUDA 思路):

image.png

在 Triton 中,可以在加载 Q、K 的切片后,应用旋转:

@triton.jit
def rope_rotate(q, k, cos, sin, seq_len, ...):
    idx = tl.program_id(0)
    # 加载 cos, sin 对应位置
    c = tl.load(cos + ...)
    s = tl.load(sin + ...)
    q0, q1 = ...
    q_rot0 = q0 * c - q1 * s
    q_rot1 = q0 * s + q1 * c
    tl.store(...)

融合后的优点:减少了一次全局内存的读写(旋转后的 Q、K 不必写出再读入),显著提升带宽利用率。实际实现较复杂,通常集成在 FlashAttention kernel 中。


基频 10000 的选取对训练稳定性的影响

基频 base 决定了频率的尺度:θ_i = base^{-2i/d}。若 base 太小(如 10),所有频率都偏大,低频分量也会快速振荡,导致模型难以捕捉长距离依赖;同时梯度在高频旋转角度上可能振荡剧烈,训练不稳定。

若 base 太大(如 10^6),则几乎所有频率都非常小,旋转角度随位置变化极慢,模型无法区分相邻位置,丧失了位置编码的意义;且对于有限训练长度,角度变化量极小,优化可能陷入平坦区域,梯度消失。

梯度分析:RoPE 通过 cos(m * θ_i)sin(m * θ_i) 作用于 Q、K。对 θ_i 的梯度包含 m 因子,因此长序列训练时,低频(i 小)的梯度较大,可能导致优化震荡;高频梯度很小,可能难以学习。适当的 base 可以平衡不同频率分量的梯度分布,使得模型能同时学习短距和长距依赖。

标准选择 base=10000 经过大量实验验证,在典型维度(如 64~128)和训练长度下提供了良好的频率覆盖和训练稳定性。当需要外推时,才通过增大 base 或 NTK-aware 方法调整。