RoPE 做什么、不做什么
对每个二维通道对,位置 的旋转矩阵:
Query/Key 变为 。因为旋转正交,保持向量范数;点积:
因此同时平移绝对位置不改变点积,只依赖相对位移 。RoPE 不直接应用于 Value,也不替代 causal mask;模型仍必须限制 Query 只看允许位置。
与 KV Cache 的位置边界
KV Cache 中历史 Keys 已使用它们原始绝对 positions 旋转并缓存。prompt 长 时,第一个输出 token 被送入下一 decode 前向,其 Query/Key 应使用 position ,不能每步把新 token position 重置为 0。position id、cache length、sliding window offset 或 prefix reuse 任一错位都会改变 attention score。
Prefix Cache 只有在 position semantics 兼容时可共享:相同 token block 若在不同 absolute offset 下使用不同 RoPE angles,KV 通常不能直接复用,除非实现有明确重定位机制。
可运行不变量与 position offset 实验
import platform,torch
torch.manual_seed(149)
DEVICE=torch.device("cuda" if torch.cuda.is_available() else "cpu")
DTYPE=torch.float64
def rotate_half(x):
x1,x2=x[...,0::2],x[...,1::2]
return torch.stack((-x2,x1),-1).flatten(-2)
def rope(x,positions,base=10_000.,scale=1.):
d=x.shape[-1]
inv=base**(-torch.arange(0,d,2,device=x.device,dtype=x.dtype)/d)
angles=(positions.to(x.dtype)/scale)[...,None]*inv
cos=torch.repeat_interleave(angles.cos(),2,-1)
sin=torch.repeat_interleave(angles.sin(),2,-1)
return x*cos+rotate_half(x)*sin
def main():
d=8; q=torch.randn(d,device=DEVICE,dtype=DTYPE); k=torch.randn_like(q)
m,n,shift=3,11,7
qm,kn=rope(q,torch.tensor(m,device=DEVICE)),rope(k,torch.tensor(n,device=DEVICE))
qs,ks=rope(q,torch.tensor(m+shift,device=DEVICE)),rope(k,torch.tensor(n+shift,device=DEVICE))
torch.testing.assert_close(qm.norm(),q.norm(),rtol=1e-12,atol=1e-12)
torch.testing.assert_close(qm@kn,qs@ks,rtol=1e-12,atol=1e-12)
t=5; keys=torch.randn(t,d,device=DEVICE,dtype=DTYPE)
cached=rope(keys,torch.arange(t,device=DEVICE))
correct=rope(q,torch.tensor(t,device=DEVICE))@cached.T
wrong=rope(q,torch.tensor(0,device=DEVICE))@cached.T
assert not torch.allclose(correct,wrong)
interpolated=rope(q,torch.tensor(16,device=DEVICE),scale=2.)
position8=rope(q,torch.tensor(8,device=DEVICE),scale=1.)
torch.testing.assert_close(interpolated,position8,rtol=1e-12,atol=1e-12)
print(f"python={platform.python_version()} torch={torch.__version__}")
print(f"device={DEVICE} dtype={DTYPE} dim={d}")
print(f"norm_error={(qm.norm()-q.norm()).abs():.3e}")
print(f"relative_shift_dot_error={((qm@kn)-(qs@ks)).abs():.3e}")
print(f"cached_K={tuple(cached.shape)} new_query_position={t}")
print(f"wrong_position_score_max_diff={(correct-wrong).abs().max():.6f}")
print(f"position_interpolation_16_scale2_equals_position8={torch.allclose(interpolated,position8)}")
if __name__=="__main__":
with torch.inference_mode(): main()
完整文件位于 examples/rope_reference.py。实际范数误差 0、同时平移点积误差 4.441e-16;把新 Query 位置从正确的 5 错写为 0,score 最大差 1.562070。Position Interpolation 的教学 position/scale 映射使 position 16、scale 2 与原 position 8 angle 相同。
用统一的 frequency schedule 比较三类扩展
令 rotary dimension 为 ,频率对索引 。标准 RoPE inverse frequency 与位置 的角度为:
常称 rope_theta 或 base。任何缩放方法最终都必须给出每个通道的 ,以及可选的 attention amplitude factor;“把最大长度从 4K 改成 16K”本身没有定义这些数值。
Position Interpolation:所有频率等比例压缩
训练长度为 ,目标长度 ,。PI 把位置映射为 :
所以实现可以除 position ids,也可以令 ;两者代数等价。目标位置 映射回训练位置 。代价是所有频率都被同样压缩,包括训练范围内原本熟悉的局部位置模式;原论文用少量 long-context finetuning 恢复质量,而不是声称零训练修改即可无损扩展。
Dynamic NTK-aware:改变 base,缩放量随维度变化
“NTK-aware”在不同代码库中不是唯一配置名。下式采用当前 Transformers 中 dynamic NTK 形式。原训练长度 、配置 factor 、当前 sequence length 下:
当 时通常保持原始 base;超过后才更新。 的最高频恒为 ,较低频被逐渐压缩,而不是 PI 的全通道 。static、dynamic、NTK-by-parts 等变体的公式和更新时间不同,因此生产配置必须记录具体 rope type、factor、original length、current/bucket length 和实现版本。
YaRN:按频段混合 extrapolation 与 interpolation
YaRN 不把全部通道一刀切。用教学抽象表示:
由某频率在原训练长度内完成的旋转次数决定,并在 beta_fast、beta_slow 对应的 correction range 中线性过渡:部分通道保持 extrapolation,部分完全 interpolation,中间通道混合。常见实现还返回 attention factor ,把生成的 cos/sin 都乘 ;于是 Q/K 范数各缩放 ,两者点积幅度会带 。具体默认值、截断规则和配置字段属于实现契约,不能只从 “YaRN” 名称猜测。
一个 的具体频率代入
取 、、、、。选择 32 个频率对中的索引 [0,1,8,16,24,31],下表列出 :
| 方法 | pair 0 | pair 1 | pair 8 | pair 16 | pair 24 | pair 31 |
|---|---|---|---|---|---|---|
| PI / linear | 0.2500 | 0.2500 | 0.2500 | 0.2500 | 0.2500 | 0.2500 |
| dynamic NTK | 1.0000 | 0.9206 | 0.5159 | 0.2661 | 0.1373 | 0.0769 |
YaRN,beta_fast=32,beta_slow=1 | 1.0000 | 1.0000 | 1.0000 | 0.6538 | 0.2500 | 0.2500 |
这张表只描述 position→angle schedule,不是质量排名。不同模型的 rotary dimension、partial rotary、训练长度、base、finetuning recipe 与 method variant 会改变结果;不能把本例参数移植为通用最优配置。
可运行的三方法与 KV 一致性实验
实验问题:PI、dynamic NTK、YaRN 的 per-frequency scaling 是否真的不同?YaRN 的 extrapolation/mixed/interpolation 区域是否满足公式?prefill 缓存 K 与 decode Q 混用两种配置会怎样?
预期结果:PI 比例全部为 0.25;dynamic NTK 比例随通道变化;YaRN 同时出现 0、部分和完整 interpolation weight。相同 schedule 的旋转保持范数,附加 attention factor 后范数按该 factor 缩放;混用 PI/YaRN 的 score 必须显著不同。代码只验证公式与状态一致性,不评估语言模型质量或速度。
import math,platform
import torch
torch.manual_seed(227)
DEVICE=torch.device("cuda" if torch.cuda.is_available() else "cpu")
DTYPE=torch.float64
def default_inv_freq(dim,base=10_000.):
pairs=torch.arange(0,dim,2,device=DEVICE,dtype=DTYPE)
return base**(-pairs/dim)
def linear_pi_inv_freq(dim,factor,base=10_000.):
return default_inv_freq(dim,base)/factor
def dynamic_ntk_inv_freq(dim,factor,original_length,sequence_length,base=10_000.):
sequence_length=max(sequence_length,original_length)
scaled_base=base*(factor*sequence_length/original_length-(factor-1))**(dim/(dim-2))
return default_inv_freq(dim,scaled_base)
def yarn_inv_freq(dim,factor,original_length,base=10_000.,beta_fast=32.,beta_slow=1.):
default=default_inv_freq(dim,base); interpolated=default/factor
def correction_dim(rotations):
return dim*math.log(original_length/(rotations*2*math.pi))/(2*math.log(base))
low=max(math.floor(correction_dim(beta_fast)),0)
high=min(math.ceil(correction_dim(beta_slow)),dim-1)
if low==high: high+=1e-3
pair_index=torch.arange(dim//2,device=DEVICE,dtype=DTYPE)
interpolation_weight=((pair_index-low)/(high-low)).clamp(0,1)
frequencies=default*(1-interpolation_weight)+interpolated*interpolation_weight
attention_factor=1. if factor<=1 else .1*math.log(factor)+1.
return frequencies,attention_factor,interpolation_weight
def rotate_half(x):
return torch.stack((-x[...,1::2],x[...,0::2]),-1).flatten(-2)
def apply_rope(x,positions,inv_freq,attention_factor=1.):
angles=positions.to(DTYPE)[...,None]*inv_freq
cos=torch.repeat_interleave(angles.cos(),2,-1)*attention_factor
sin=torch.repeat_interleave(angles.sin(),2,-1)*attention_factor
return x*cos+rotate_half(x)*sin
def main():
dim,original_length,factor=64,4096,4.
target_length=int(original_length*factor)
default=default_inv_freq(dim); linear=linear_pi_inv_freq(dim,factor)
dynamic=dynamic_ntk_inv_freq(dim,factor,original_length,target_length)
yarn,attention_factor,weight=yarn_inv_freq(dim,factor,original_length)
torch.testing.assert_close(linear,default/factor,rtol=0,atol=0)
assert dynamic[0]==default[0] and torch.all(dynamic[1:]<default[1:])
assert torch.any(weight==0) and torch.any((weight>0)&(weight<1)) and torch.any(weight==1)
torch.testing.assert_close(yarn[weight==0],default[weight==0])
torch.testing.assert_close(yarn[weight==1],linear[weight==1])
vector=torch.randn(dim,device=DEVICE,dtype=DTYPE)
position=torch.tensor(target_length-1,device=DEVICE)
rotated=apply_rope(vector,position,yarn)
scaled=apply_rope(vector,position,yarn,attention_factor)
torch.testing.assert_close(rotated.norm(),vector.norm(),rtol=1e-12,atol=1e-12)
torch.testing.assert_close(scaled.norm(),vector.norm()*attention_factor,rtol=1e-12,atol=1e-12)
tokens=12; keys=torch.randn(tokens,dim,device=DEVICE,dtype=DTYPE)
query=torch.randn(dim,device=DEVICE,dtype=DTYPE)
cached_keys=apply_rope(keys,torch.arange(tokens,device=DEVICE),yarn,attention_factor)
correct_query=apply_rope(query,torch.tensor(tokens,device=DEVICE),yarn,attention_factor)
wrong_query=apply_rope(query,torch.tensor(tokens,device=DEVICE),linear)
correct_scores=correct_query@cached_keys.T/math.sqrt(dim)
wrong_scores=wrong_query@cached_keys.T/math.sqrt(dim)
mismatch=(correct_scores-wrong_scores).abs().max(); assert mismatch>1e-3
selected=torch.tensor([0,1,8,16,24,31],device=DEVICE)
ratios={"linear":(linear/default)[selected].cpu().tolist(),
"dynamic_ntk":(dynamic/default)[selected].cpu().tolist(),
"yarn":(yarn/default)[selected].cpu().tolist()}
print(f"python={platform.python_version()} torch={torch.__version__}")
print(f"device={DEVICE} dtype={DTYPE} rotary_dim={dim} frequency_pairs={dim//2}")
print(f"original_length={original_length} target_length={target_length} factor={factor}")
print(f"selected_pair_indices={selected.cpu().tolist()} frequency_ratios={ratios}")
print(f"yarn_interpolation_weights={weight[selected].cpu().tolist()}")
print(f"yarn_attention_factor={attention_factor:.6f} norm_ratio={(scaled.norm()/vector.norm()):.6f}")
print(f"prefill_K={tuple(cached_keys.shape)} decode_Q={tuple(correct_query.shape)}")
print(f"mixed_config_score_max_abs_diff={mismatch:.6f}")
if __name__=="__main__":
with torch.inference_mode(): main()
完整文件位于 examples/rope_scaling_methods.py。2026-08-13 在 PyTorch 2.13.0+cu130 上,输出得到上表的 frequency ratios;YaRN interpolation weights 为 [0,0,0,0.4615,1,1],attention factor 与范数比均为 1.138629。缓存 K 用 YaRN、decode Q 错用 PI 时,score 最大绝对差为 2.490139。
教学实现遵循主流 Transformers 当前公式,但没有依赖该包,也没有声称所有模型 config 使用同一 variant。它用 FP64 隔离公式误差;生产中 sin/cos cache、低精度、partial rotary、fused QK kernel 和 KV layout 都会影响实现证据。
计算、显存、延迟与质量成本
缩放方法本身通常只改变长度为 的 frequency table、cos/sin 生成或 fused kernel 内角度计算,metadata 很小。真正昂贵的是它允许请求进入更长的 :
- full causal prefill attention 的逻辑 score 工作随 增长;FlashAttention可减少中间 IO,不改变 dense 依赖数;
- KV Cache payload 随 增长,decode 每步读取历史也随可见 增长;
- sin/cos table 若预计算到最大长度,容量约为
T × rotary_dim × element_size(是否同时存 cos/sin、是否广播/分层共享依实现而定);按需生成则转移为计算与 kernel/fusion 约束; - dynamic schedule 可能引入 shape/length guards、重新生成 tables 或编译 bucket;若改变已缓存 token 的频率语义,成本可能升级为 KV 重算;
- attention factor、频率压缩和超出训练分布会改变 logits/质量,必须用 perplexity、long-context retrieval、真实任务与位置分桶评估,而不是只看代码能否处理更长 tensor。
这些方法不减少 。若目标是限制可见历史和 KV 上限,应看 Sliding Window Attention;若单卡装不下完整上下文,可看 Context Parallel,但它加入通信且仍需全局 softmax 语义。
常见失效模式
- prefill 与 decode 使用不同 scaling/base 或 position offset;
- dynamic NTK 随 sequence length 更新 frequency schedule,却继续读取按旧 schedule 旋转的缓存 K;
- prefix cache 在不同 RoPE 配置/absolute offsets 间误复用;
- sliding window 丢弃旧 KV,却仍按错误相对/绝对位置旋转;
- 低精度 sin/cos 或超大 positions 带来相位精度问题;
- backend 使用预计算 cos/sin cache,扩长后未正确扩容或索引;
- 用某框架的
linear、dynamic、yarn配置字段套到另一版本/模型,却没有核对公式、original length、attention factor 与 partial rotary; - 只测 “needle” 命中,不测长序列困惑度、真实任务、不同信息位置和延迟/容量。
上下游关系
- Tokenizer/Chat Template决定实际 token positions 与长度。
- KV Cache缓存已旋转 Keys;decode position 必须接续 cache length。
- Sliding Window Attention限制可见历史并可截断 cache;RoPE scaling 修改角度映射。二者可组合但解决的问题不同。
- Attention Sink固定保留 early tokens 并滚动 recent KV;若 runtime 把 recent Keys 重映射到 local positions,必须按本页 frequency schedule 正确 rerotate,而不是让物理 slot 决定位置。
- Sequence/Context Parallel可把长 context 的 KV 所有权分到多个 rank;它解决容量/计算放置,不定义本页 frequency schedule。
- Prefix Cache identity 要包含 RoPE/scaling/offset 语义。
- MLA把位置 key 解耦成单独路径,但仍需正确 RoPE positions。
- FlashAttention和 PagedAttention优化 IO/布局,不自动提升长上下文模型质量。
参考资料
- Su et al., RoFormer: Enhanced Transformer with Rotary Position Embedding
- Chen et al., Extending Context Window of Large Language Models via Positional Interpolation
- Peng et al., YaRN: Efficient Context Window Extension of Large Language Models
- Hugging Face Transformers — RoPE utilities
- Hugging Face Transformers —
modeling_rope_utils.py— linear、dynamic NTK、YaRN 的当前实现公式与配置字段。