NTK-aware Scaling、YaRN 等方法让 RoPE 模型处理超长序列
配套代码
假设模型在 max_seq_len = 4096 上训练。当推理时输入长度为 8192,会发生什么?
一个常见的直觉是”位置编号超出了训练范围,所以 OOD 了”。但这个说法不够精确,位置编号跟位置嵌入不是一回事。位置编号 𝑚 m 是无界的,但 RoPE 的位置嵌入是三角函数组成的,有界。跟模型直接打交道的是位置嵌入,不是位置编号。所以要真正理解 OOD,我们需要从位置嵌入的角度来分析。
回顾上一章的内积公式,加了 RoPE 之后的 Q/K 内积可以用复数表示为:
cos 𝑡 + 𝑖 sin 𝑡 e it =cost+isint 可以知道,它就是单位圆上的一个点。当相对距离 𝑚 − 𝑛 m−n 逐渐变大时,这个点在单位圆上转圈,𝜃 𝑖 θ i 越大转得越快,𝜃 𝑖 θ i 越小转得越慢。
这就是”转圈视角”的核心:位置编号 𝑚 − 𝑛 m−n 是否 OOD 根本不重要,重要的是单位圆上的点是否被充分训练过。
假设训练长度为 𝐿 train L train ,那么 𝑚 − 𝑛 ∈ [ 0 , 𝐿 train − 1] m−n∈[0,L train
−1]。对于每个维度 𝑖 i,我们可以算出训练期间转了多少圈:
高频维度( 𝜃 𝑖 θ i
大, 𝑖 i 小):转速快,训练期间已经转了很多圈,圆上的每一个点几乎都被训练过。即使测试时 𝑚 − 𝑛 m−n 更大,也只是在已经见过的圆上继续转,不存在 OOD 问题
低频维度( 𝜃 𝑖 θ i
小, 𝑖 i 大):转速慢,训练期间可能还没转完一圈,被训练过的点顶多只是圆上的一段弧。测试时如果超出了这段弧的范围,就进入了模型从未见过的区域,这才是真正的 OOD
Unit Circle Coverage
L_train
4096
L_test
8192