全部术语

GLOSSARY ENTRY

RoPE、Position Interpolation、NTK-aware 与 YaRN

  • Rotary Position Embedding
  • RoPE
  • Position Interpolation
  • Dynamic NTK Scaling
  • NTK-aware Scaling
  • YaRN
  • RoPE Scaling

RoPE 按位置旋转 Query/Key 的二维通道对;PI、NTK-aware 与 YaRN 用不同频率计划扩展位置范围,但必须与训练分布、prefill/decode 的 KV 状态和 backend 配置一致。

RoPE 做什么、不做什么

对每个二维通道对,位置 mm 的旋转矩阵:

Rm(θ)=[cos(mθ)sin(mθ)sin(mθ)cos(mθ)]R_m(\theta)=\begin{bmatrix} \cos(m\theta)&-\sin(m\theta)\\ \sin(m\theta)&\cos(m\theta) \end{bmatrix}

Query/Key 变为 Rmq,RnkR_mq,R_nk。因为旋转正交,保持向量范数;点积:

(Rmq)T(Rnk)=qTRnmk(R_mq)^\mathsf T(R_nk)=q^\mathsf TR_{n-m}k

因此同时平移绝对位置不改变点积,只依赖相对位移 nmn-m。RoPE 不直接应用于 Value,也不替代 causal mask;模型仍必须限制 Query 只看允许位置。

与 KV Cache 的位置边界

KV Cache 中历史 Keys 已使用它们原始绝对 positions 旋转并缓存。prompt 长 TT 时,第一个输出 token 被送入下一 decode 前向,其 Query/Key 应使用 position TT,不能每步把新 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 为 dd,频率对索引 i=0,,d/21i=0,\ldots,d/2-1。标准 RoPE inverse frequency 与位置 mm 的角度为:

ωi=θ2i/d,ϕi(m)=mωi\omega_i=\theta^{-2i/d},\qquad \phi_i(m)=m\omega_i

θ\theta 常称 rope_theta 或 base。任何缩放方法最终都必须给出每个通道的 ωi\omega'_i,以及可选的 attention amplitude factor;“把最大长度从 4K 改成 16K”本身没有定义这些数值。

Position Interpolation:所有频率等比例压缩

训练长度为 LL,目标长度 L=fLL'=fLf>1f>1。PI 把位置映射为 m=m/fm'=m/f

ϕiPI(m)=mfωi=mωif\phi_i^{PI}(m)=\frac{m}{f}\omega_i=m\frac{\omega_i}{f}

所以实现可以除 position ids,也可以令 ωiPI=ωi/f\omega_i^{PI}=\omega_i/f;两者代数等价。目标位置 m=Lm=L' 映射回训练位置 LL。代价是所有频率都被同样压缩,包括训练范围内原本熟悉的局部位置模式;原论文用少量 long-context finetuning 恢复质量,而不是声称零训练修改即可无损扩展。

Dynamic NTK-aware:改变 base,缩放量随维度变化

“NTK-aware”在不同代码库中不是唯一配置名。下式采用当前 Transformers 中 dynamic NTK 形式。原训练长度 LL、配置 factor ff、当前 sequence length TT 下:

θ=θ(fTL(f1))d/(d2),ωiNTK=(θ)2i/d\theta'=\theta\left(f\frac{T}{L}-(f-1)\right)^{d/(d-2)}, \qquad \omega_i^{NTK}=(\theta')^{-2i/d}

TLT\le L 时通常保持原始 base;超过后才更新。i=0i=0 的最高频恒为 11,较低频被逐渐压缩,而不是 PI 的全通道 1/f1/fstaticdynamicNTK-by-parts 等变体的公式和更新时间不同,因此生产配置必须记录具体 rope type、factor、original length、current/bucket length 和实现版本。

YaRN:按频段混合 extrapolation 与 interpolation

YaRN 不把全部通道一刀切。用教学抽象表示:

ωiextra=ωi,qquadωiinterp=ωif\omega_i^{extra}=\omega_i,qquad \omega_i^{interp}=\frac{\omega_i}{f} ωiYaRN=(1ri)ωiextra+riωiinterp,0ri1\omega_i^{YaRN}=(1-r_i)\omega_i^{extra}+r_i\omega_i^{interp}, \qquad 0\le r_i\le1

rir_i 由某频率在原训练长度内完成的旋转次数决定,并在 beta_fastbeta_slow 对应的 correction range 中线性过渡:部分通道保持 extrapolation,部分完全 interpolation,中间通道混合。常见实现还返回 attention factor aa,把生成的 cos/sin 都乘 aa;于是 Q/K 范数各缩放 aa,两者点积幅度会带 a2a^2。具体默认值、截断规则和配置字段属于实现契约,不能只从 “YaRN” 名称猜测。

一个 4K16K4K\rightarrow16K 的具体频率代入

d=64d=64θ=10000\theta=10000L=4096L=4096f=4f=4T=16384T=16384。选择 32 个频率对中的索引 [0,1,8,16,24,31],下表列出 ωi/ωi\omega'_i/\omega_i

方法pair 0pair 1pair 8pair 16pair 24pair 31
PI / linear0.25000.25000.25000.25000.25000.2500
dynamic NTK1.00000.92060.51590.26610.13730.0769
YaRN,beta_fast=32,beta_slow=11.00001.00001.00000.65380.25000.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 都会影响实现证据。

计算、显存、延迟与质量成本

缩放方法本身通常只改变长度为 d/2d/2 的 frequency table、cos/sin 生成或 fused kernel 内角度计算,metadata 很小。真正昂贵的是它允许请求进入更长的 TT

  • full causal prefill attention 的逻辑 score 工作随 O(T2)O(T^2) 增长;FlashAttention可减少中间 IO,不改变 dense 依赖数;
  • KV Cache payload 随 O(T)O(T) 增长,decode 每步读取历史也随可见 TT 增长;
  • 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。

这些方法不减少 TT。若目标是限制可见历史和 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,扩长后未正确扩容或索引;
  • 用某框架的 lineardynamicyarn 配置字段套到另一版本/模型,却没有核对公式、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。
  • FlashAttentionPagedAttention优化 IO/布局,不自动提升长上下文模型质量。

参考资料