大模型在训练的时候,句子里的每个词都不是孤立存在的。它们需要知道彼此的位置:谁在前面,谁在后面,隔了多远。Transformer 结构本身是“一视同仁”的,它不会自动记录顺序,所以研究者发明了很多位置编码方法。RoPE,全称 Rotary Position Embedding,也就是旋转位置编码,是目前应用非常广的一种。它没有用一张表格去硬记每个位置,而是把位置信息藏在向量的旋转角度里,让每个 token 的位置像一根表针一样,转一下就知道自己站在哪。
一、为什么需要位置编码
1.1 Transformer 不记事
你往模型里扔一句话,模型看到的其实是一个一个向量组成的序列。计算注意力时,模型会对当前向量和序列里其他向量做点积,然后看谁和自己更像。问题是,点积本身不带顺序,比如“我给你钱”和“你借钱我”这两个句子,词都差不多,但意思完全不一样。Transformer 如果不知道位置,就会觉得这两个句子没什么区别。
所以必须有人工加入的位置信号。位置编码的任务,是告诉模型“我是第几个词”,或者更进一步,“我离你有多远”。
1.2 老办法有什么问题
最早常用的方式是绝对位置编码,比如训练一个巨大的矩阵,按位置 index 去查一张表。这种做法简单粗暴,但有两个毛病:一是占内存,序列越长,表越大;二是不会举一反三,训练时见过 2048 个位置,突然让它处理 4096 个位置,它就没见过,容易懵。
后来有人用正弦余弦函数生成固定位置信号,虽然不用学参数,但表达相对位置关系的能力还是不够直接。RoPE 则想了一个更聪明的办法:用旋转角度表示位置,让 query 和 key 之间的距离,天然体现在注意力分数里。
二、RoPE 的核心直觉
2.1 把位置当成旋转角度
你想象一个钟表,指针每走一格,角度就变一点。RoPE 也是这样,它把一个向量旋转一个角度,这个角度由位置编号 m 控制。位置 0 不转,位置 1 转一点,位置 2 转更多一点,每个位置都有自己的角度状态。
在二维平面上,一个向量 (x0, x1) 旋转角度 θ 后,会变成 (x0 * cosθ - x1 * sinθ, x0 * sinθ + x1 * cosθ)。这个公式看起来复杂,其实就是在画圆。把位置 m 编进 θ,也就把位置信息写进了向量本身。
2.2 高维向量怎么办
大模型里每个 head 的维度不是 2,而是 64、96 或 128。RoPE 的做法是:把向量拆成很多个二维小对,比如第 0 和第 1 维是一对,第 2 和第 3 维是一对,依此类推。每一对都旋转,但旋转的“速度”不一样。
你给位置 m 一个频率 θi,不同维度的频率不同。低维部分转得快,敏感地捕捉相邻 token 的位置;高维部分转得慢,负责捕捉长距离信息。这样整个向量经过一次旋转,既知道自己附近有什么,也知道自己离远处的东西多远。
2.3 为什么它能表示相对位置
RoPE 有个特别妙的性质:两个向量如果都旋转了各自的位置角度,再把它们做点积,结果只和角度差有关,也就是和两个 token 之间的相对距离有关。这就是大家常说的“旋转后点积自带相对位置信息”。对语言模型来说,它不需要知道绝对的第几号位,更关心“这个词和前面第三个词是不是有关联”,RoPE 天然符合这个需求。
三、用 PyTorch 实现和接入
下面的代码统一使用 Python 3.10 和 PyTorch 2.0,所有示例都是可运行的最小版本,方便你拆开看每一行在做什么。
3.1 预计算频率表
# 技术栈:Python 3.10 + PyTorch 2.0
import torch
def precompute_freqs(head_dim, max_seq_len, base=10000.0):
# head_dim 必须能被2整除,比如64、96、128
# 计算每个旋转对儿的频率,这里取偶数下标,数量是 head_dim/2
inv_freq = 1.0 / (
base ** (torch.arange(0, head_dim, 2).float() / head_dim)
)
# 生成位置编号:0,1,2,... max_seq_len-1
positions = torch.arange(max_seq_len, dtype=torch.float32)
# 外积:positions 的每一行,乘上 inv_freq 的每一个频率
# 得到形状 (max_seq_len, head_dim/2)
freqs = torch.einsum("m,i->mi", positions, inv_freq)
# 把 freqs 复制一份拼起来。
# 以 head_dim=4 为例,freqs 原本是 [theta0, theta1],
# cat 之后变成 [theta0, theta1, theta0, theta1]。
# 这样 cos/sin 可以直接和完整的 q/k 向量按维度相乘。
freqs_2d = torch.cat([freqs, freqs], dim=-1)
cos = freqs_2d.cos()
sin = freqs_2d.sin()
return cos, sin
这段代码看起来短,但它是整个 RoPE 的地基。cos 和 sin 会按照位置生成好,后面每个 token 都在这个“角度表”里找到自己的位置。
3.2 实现旋转 Q 和 K
有了 cos 和 sin,接下来要做的是把 query 和 key 旋转一下。注意,value 是不旋转的,因为位置信息只需要影响注意力分数的计算,不需要污染最终输出的内容。
def rotate_half(x):
# 把 x 分成前后两半,然后交换位置并给前半部分加负号
half = x.shape[-1] // 2
x1 = x[..., :half] # 前半段
x2 = x[..., half:] # 后半段
return torch.cat([-x2, x1], dim=-1)
def apply_rotary_emb(q, k, cos, sin):
# cos/sin 的形状是 (seq_len, head_dim)
# 加上 batch 和 head 两个维度,方便和 q/k 广播
# q/k 的形状是 (batch, heads, seq_len, head_dim)
cos = cos.unsqueeze(0).unsqueeze(0) # 变成 (1,1,seq_len,head_dim)
sin = sin.unsqueeze(0).unsqueeze(0)
# 旋转公式:新向量 = 原向量 * cos + rotate_half(原向量) * sin
q_embed = q * cos + rotate_half(q) * sin
k_embed = k * cos + rotate_half(k) * sin
return q_embed, k_embed
这段代码里最关键的是 rotate_half。它看起来只是在拼接两个张量,实际上它模拟了复数乘法里的虚数部分。理解了它,就理解了“旋转”是怎么落到张量计算上的。
3.3 在注意力层里用起来
实战中,RoPE 不是单独用的,而是插在注意力层里:从输入算出 q、k、v 之后,先旋转 q、k,再做注意力运算。
# 技术栈:Python 3.10 + PyTorch 2.0
import torch.nn as nn
import math
class LlamaAttentionRoPE(nn.Module):
def __init__(self, hidden_dim, n_heads, max_seq_len):
super().__init__()
self.n_heads = n_heads
self.head_dim = hidden_dim // n_heads
# 三个线性层分别生成 q、k、v
self.wq = nn.Linear(hidden_dim, hidden_dim)
self.wk = nn.Linear(hidden_dim, hidden_dim)
self.wv = nn.Linear(hidden_dim, hidden_dim)
# 输出映射
self.wo = nn.Linear(hidden_dim, hidden_dim)
# 预计算好 max_seq_len 以内的 cos 和 sin
self.cos, self.sin = precompute_freqs(self.head_dim, max_seq_len)
def forward(self, x):
batch, seq_len, _ = x.shape
# 生成 q/k/v,并拆成多头
# x 从 (batch, seq_len, hidden_dim) 变成 (batch, heads, seq_len, head_dim)
q = self.wq(x).view(batch, seq_len, self.n_heads, self.head_dim).transpose(1, 2)
k = self.wk(x).view(batch, seq_len, self.n_heads, self.head_dim).transpose(1, 2)
v = self.wv(x).view(batch, seq_len, self.n_heads, self.head_dim).transpose(1, 2)
# 取当前 seq_len 对应的 cos/sin
cos = self.cos[:seq_len].unsqueeze(0).unsqueeze(0)
sin = self.sin[:seq_len].unsqueeze(0).unsqueeze(0)
# 只旋转 q 和 k,v 不旋转
q, k = apply_rotary_emb(q, k, cos, sin)
# 标准缩放点积注意力
scores = q @ k.transpose(-2, -1) / math.sqrt(self.head_dim)
scores = scores.softmax(dim=-1)
out = scores @ v
out = out.transpose(1, 2).contiguous().view(batch, seq_len, -1)
return self.wo(out)
从模型角度看,RoPE 没有引入新的可训练参数。它只是把 q、k 在进入注意力计算之前“拧”了一下,剩下的训练过程几乎不用改动。这也是它能被 LLaMA 等模型大刀阔斧使用的原因之一。
3.4 训练高频执行时要记得缓存
大模型训练时,每步都要处理成千上万个 token。如果每次 forward 都重新生成 cos 和 sin,虽然计算量不大,但白白浪费显存和算力。更常见的做法是启动时预计算好,训练过程中直接从缓存里取。
# 技术栈:Python 3.10 + PyTorch 2.0
_rope_cache = {}
def get_rope_cache(head_dim, seq_len, base=10000.0):
# 用 head_dim 和 seq_len 和 base 做 key,避免重复计算
key = (head_dim, seq_len, base)
if key not in _rope_cache:
cos, sin = precompute_freqs(head_dim, seq_len, base)
_rope_cache[key] = (cos, sin)
return _rope_cache[key]
很多开源的 LLaMA 实现里都有类似缓存。别小看这个优化,在大规模训练中,每个小优化都可能省出几卡 GPU。
四、大规模训练里的实际应用
4.1 长文本训练时的显存控制
RoPE 有一点很友好:它不需要额外维护一张大的位置 embedding 表。LLaMA 这类模型在预训练时经常要把 batch 拉大,显存非常紧张。如果用绝对位置编码,训练序列从 2048 扩到 8192,位置表也会跟着变大。但 RoPE 只有 cos 和 sin 两组张量,而且它们是按需生成的,不参与梯度更新,所以显存开销小得多。
4.2 上下文窗口扩展离不开 RoPE 的改造
RoPE 虽然很香,但并不是说它天然就能支持无限长的序列。如果你把一个只训练到 2048 的模型直接扔到 4096 的文本上,注意力分数会变得不可控,模型表现可能崩掉。于是大规模训练里出现了一类实际操作:把 base 从 10000 调大,或者用 NTK 缩放等方式改变旋转频率,让低维部分转得慢一点,适应更长的上下文。这种做法的核心思路是:频率变了,位置信息仍然存在,只是“刻度”被拉伸了。
# 技术栈:Python 3.10 + PyTorch 2.0
def precompute_freqs_scaled(head_dim, max_seq_len, scale=1.0):
# scale 大于1时,base 变大,低频部分旋转速度变慢
# 相当于把原本的“角度标尺”拉长,能表达更远的位置
base = 10000.0 * scale
inv_freq = 1.0 / (
base ** (torch.arange(0, head_dim, 2).float() / head_dim)
)
positions = torch.arange(max_seq_len, dtype=torch.float32)
freqs = torch.einsum("m,i->mi", positions, inv_freq)
freqs_2d = torch.cat([freqs, freqs], dim=-1)
return freqs_2d.cos(), freqs_2d.sin()
这只是演示基础思想,真正的 NTK 动态缩放还会根据当前序列长度动态调整 scale,但看懂这个例子后,再去看那些花哨的实现都会容易很多。
4.3 和 KV Cache 配合
在自回归生成时,模型是一个 token 一个 token 往出蹦的。为了省算力,模型会把历史 token 的 key 和 value 缓存起来,这就是常说的 KV Cache。RoPE 和 KV Cache 配合时需要注意的是:k 在进入缓存前就要完成旋转,这样后面新的 token 来算注意力时,直接用旋转后的历史 k 和当前旋转后的 q 做点积就行,不需要把历史 k 再旋转一次。很多新手在这里容易写错,导致训练时没问题,推理时结果突然变差。
五、优点、缺点和注意事项
5.1 优点
RoPE 最大的优点是“用旋转表达位置”这个思路本身很优雅。它不需要学习位置参数,省显存;它能让注意力分数只依赖相对距离,泛化性更好;它还可以配合各种缩放方法,在长文本场景被“抢救”回来。再加上实现足够简单,只需要十几行代码,就能嵌入到各种 Transformer 结构里,所以 LLaMA、Mistral、Qwen 这些模型都对它情有独钟。
5.2 缺点
RoPE 也不是没有短板。第一,它天生对“外推”不友好,训练长度之外的位置,模型很容易懵。第二,实现上容易踩坑,尤其是维度配对和广播,一旦出错,模型还能训练,但效果会很差,并且很难排查。第三,它只作用于 q 和 k,对 v 无能为力,所以位置信息是通过注意力权重传播的,并不是直接融合进所有向量里,这本身也限制了它表达位置信息的完整度。
5.3 注意事项
如果你要在自己的模型里接 RoPE,有几件事要盯死:head_dim 一定要是偶数,否则没法配对;q 和 k 必须使用同一套 cos 和 sin;v 不要旋转;cos 和 sin 的序列长度一定要大于等于实际序列长度,否则会越界;在混合精度训练里,建议保留 float32 的 cos 和 sin,再在计算注意力时转成 bf16 或 fp16,精度会更稳。另一个容易被忽视的点是,NLP 里很多人都喜欢用单一模型结构做“炼丹”,但不同模型的 head_dim 不一样,换模型时要把频率表跟着重算一遍。
六、总结
RoPE 是 LLaMA 这类模型里非常关键的一块拼图。它把一个看似高深的“位置编码”问题,变成了“向量旋转”问题。你不需要背公式,只要记住钟表盘上指针转动的感觉就可以了:每个 token 都有一个属于自己的角度,token 之间离得越远,角度差越大,注意力分数上的惩罚也就越明显。在大规模训练中,RoPE 靠省内存、好接入、可扩展这几个优点站稳了脚跟,但它也有外推困难等实际问题,需要靠缓存、NTK 缩放、合理的精度控制来弥补。把这套原理和代码都吃透,你再去读任何一代 LLaMA 的源码,都会觉得思路通透明朗。
评论
围绕“LLaMA模型架构中RoPE位置编码的实现原理及其在大规模训练中的实际应用”参与讨论