位置编码 入门
注意力把序列当集合读:把 token 洗一遍牌,每个输出原样回来。用你能亲手驱动的图证明三件事 —— 好编码必须有界;一张正弦频率表让点积只依赖于间隔;而把 query 和 key 转起来,这一点就从提示变成分数本身的不变式。
注意力看不见任何东西在哪儿
把 token 洗一遍牌,每个输出原封不动地回来,只是换了个排法。补救办法是一次加法 —— 而这次加法并不免费。
看看分数里到底有什么。Q[i]·K[j] 读的是两个 token 的内容,别的什么也没读;softmax(QKᵀ/√d)·V 里任何地方都没有下标。注意力具有置换等变性:输入行打乱,输出行跟着打乱。
下面是三个 token 的六种排列,每个 token 下面画着它的注意力输出。拖动顺序,看着输出沿着连线跟着自己那个词一起走:
注意读数从来没离开过 0.00。「狗咬人」和「人咬狗」交给注意力的是同样的三个向量,只是换了个摆法 —— 对语言来说这就是错的答案,因为后面任何一层都没法再还原出谁是施事。
于是我们在注意力看到之前,先把位置塞进向量里。每个 token 的嵌入上都加一个只依赖于它待在哪儿的向量 —— 拖动位置,看中间那行变,最上面那行纹丝不动:
最下面那行才是注意力真正读到的东西。看 ‖PE(p)‖ 无论滑块拖到哪儿都停在 4.00:每一对槽位都是单位圆上的一个点,十六对加起来长度恒为 4。「每个位置都有界」是一套编码首先必须做对的事。
可是现在一个向量要扛两件事,它们会互相干扰。把位置权重拧大,比较两种相似度 ——同一个词相隔 5,对上两个不同的词待在同一个槽位:
两条曲线在处交叉。过了这个点,两个不同的词待在同一个槽位比同一个词相隔 5 还要像,上面每一层都得去拆一个已经丢掉了区分度的和。2017 年那篇论文用的权重是 1,稳稳地落在交叉点左边。
所以清单是:有界、每个位置各不相同、安静到内容能活下来,最好还能说点关于距离的事。下一节先试所有人第一时间想到的两种方案。
数数比看上去难
两种一眼就能想到的编码,以及各自具体是怎么崩的。活下来的,是同时在好几个尺度上数数的那种。
最简单的编码就是下标本身:位置 0 拿 0,位置 1 拿 1,一直下去。不用训练,不用存,也没有最大长度 —— 同时也是死得最快的那个。
它要加上去的那个向量,长度大约是 1。把标记沿着斜坡往上走,看它多快就冲出了嵌入所在的那条带子:
到的时候,编码已经是它所加之物的六十三倍。读这个和的那一层几乎只看得见位置,回传到词嵌入上的梯度被彻底淹没。有界不是锦上添花的要求。
显而易见的补救是除以序列长度,这确实把它压住了。把第二条序列拉长,看槽位 3 在两边各值多少:
注意槽位 3 在8 个 token 的序列里是 0.43,在 64 个 token 的序列里是 0.05。编码不再是位置的属性,而变成了这一批数据的属性 ——「往回数三个 token」从此没有一个固定的表示可供模型学习。
我们要的是:有界、绝对,并且仍然携带距离信息。二进制三样都占齐了:拖动位置,看每一位按自己的节奏翻转:
第 0 位每 2 个 token 翻一次,第 3 位每 16 个翻一次。从上往下读:快的位分辨相邻,慢的位分辨整片区域,而且不管序列多长,每个值都待在 {0, 1} 里。
这台里程表唯一的毛病是它由台阶组成:01111 走到 10000 一次翻掉五位,梯度找不到平滑方向。把拐角磨圆,就是正弦编码。
一台磨圆了角的里程表
正弦和余弦,架在一把等比的频率梯子上。没有参数,每个尺度都配一个波长。
2017 年的公式把位置 p 的第 2i 个槽位定成 sin(p·θᵢ),第 2i+1 个定成 cos(p·θᵢ),其中 θᵢ = base^(−2i/d)。每一对槽位分到自己的频率,按等比递减。
这就是把拐角磨圆之后的里程表。挑一对出来,从读数里读它的频率 —— 上面几行几个 token 就转完一圈,下面几行几乎不弯:
第 0 对每个 token 转 1.000 弧度,所以每 6.3 个 token 转回原处。要 353 个。梯子做成等比是有意的:每一级都只花同样的两个槽位,却换来比上一级粗一个数量级的尺度。
一个位置的编码,就是竖着把它们全部切开。滑动这一刀,看八个点各自落到自己的高度:
这八个正弦值,加上它们的八个余弦搭档,就是 PE(p) 的十六个数。没有两个位置共用同一列,而相邻的列只在快的那几行上有差别 —— 这正是二进制那套距离感,只不过现在是连续的。
梯子能伸多远由底数决定。先挑一对,再换它底下的底数:
在论文用的底数 10000 下,这把梯子从第 0 对的 6.3 个 token 一路到的 35333 个。十六级里有五级落在训练窗口之上,在里面转不满一圈。把底数调大,整把梯子跟着拉长,这正是 Llama 3 扳动的那根杆。
于是编码做到了有界、每个位置各不相同、多尺度,而且免费。剩下的是后面一切都依赖的那条性质:它到底有没有说出两个位置之间的距离。
要的是间隔,不是位子
每一对槽位都是圆上的一个点,所以往前走就是转动。正是这一点让正弦频率表不只是一个哈希。
单独拿一对出来看。(sin p·θᵢ, cos p·θᵢ) 是单位圆上角度为 p·θᵢ 的一个点;从 p 走到 p+k 就是转过 k·θᵢ —— 而这个转角跟 p 一点关系都没有。
挪动位置,看它的弧越来越长,而往前那一步始终保持同样的大小:
因为不论位置停在哪儿,那一步都是同样大的转角,所以 PE(p+k) 是 PE(p) 的一个线性函数,而这个矩阵只依赖于 k。一个权重矩阵就能对所有位置同时实现「往回看三个 token」。
在全部十六对上同时这么做,就有东西塌下来了。在波形堆上竖着切两刀,读一读它们挑出的两列之间的点积:
把两刀一起挪,这个数不动。间隔为 8 时,不论这两刀落在哪里,读数都是 0.66;只有当两刀重合时它才升到 1.00。
同一件事画成曲线。固定第一个位置,拿它去和其他每个位置打分 —— 整条曲线跟着走,形状分毫不变:
因为 sin a sin b + cos a cos b 精确等于 cos(a−b),所以 PE(p)·PE(q) 就是 Σᵢ cos((p−q)·θᵢ) —— 只跟间隔有关,跟别的一概无关。往后数八个 token 的那个值,无论 p 是 0 还是,都停在 0.66。
下面这一段是各种总结通常会跳过的。注意力算的不是 PE 对 PE,而是 (x+PE)W_Q 对 (x+PE)W_K。把投影从恒等映射上推开:
静止时八条曲线严丝合缝地压在精确曲线上,投影一动它们就散开:混合到 1.00 时,间隔为 8 处的离散度是 0.48,接近整个值域的四分之一。平移不变性是模型可以用的一组基,从来不是它一定拿得到的保证。
更糟的是,(x+PE)W_Q · (x+PE)W_K 展开有四项,只有一项是位置对位置。编码承诺的和分数交付的之间这道缝,正是 RoPE 钻进来的口子。
给每个位置分一行
BERT 和 GPT-2 干脆跳过了算术:每个位置分一行参数,里面装什么交给优化器决定。
nn.Embedding(max_len, d_model),按下标查表,然后像正弦向量一样加上去。跟词表同一条代码路径,只是换了套词汇 —— 写过词嵌入的人等于已经写过它。
没有任何东西约束一行里能放什么。挑一行出来,再拖平滑度 —— 两个极端都是梯度下降有可能留下的状态:
在上,各行就是彼此独立的随机抽样,而这正是参数化本身所保证的东西:什么都不保证。训练完的表通常比这平滑,但模型从来没有要求过这一点 —— 那份结构是数据留下的痕迹,不是方案自带的性质。
这决定了相邻两行长什么样。把学出来的相似度和同样间隔下精确的正弦相似度放在一起比:
因为慢的那几对几乎没动,正弦曲线在间隔 1 处是 0.96,然后平滑地往下走。学习曲线的起点则完全看训练把它留在哪儿 —— 这里在平滑度 0.60 上是 0.60 —— 而所有数据没有练到的位置对,都是抛硬币。
真正致命的问题比这些都简单。把序列推到表的末端外面去:
根本没有第 1024 行。GPT-2 small 为位置分配了 1024 × 768 = 786432 个参数,然后就到头了;问它,查表直接抛 IndexError。这是少见的、会大声报错的失败。
这就是那笔交易。学习表在 max_len 以内拟合得完美,过了线一个字也说不出来;公式则处处都能说点什么,却哪儿都不完美。 RoPE 保留公式,只改它作用的地方。
不加,改成转
别动嵌入。在分数内部,让 query 和 key 各自按自己的位置转一个角度,最后活下来的只有间隔。
RoPE(Su 等,2021)根本不碰词向量。在每个头内部,它把 Q 和 K 的每一对维度当成二维向量,按 §03 那张频率表转过 p·θᵢ。
圆上两个向量的点积只取决于它们之间的夹角。挪动query 的位置和key 的位置,盯着分数看:
把两个一起挪同样多,分数纹丝不动;只挪一个,它就只随间隔变化。这就是那条不变式:⟨R_m q, R_n k⟩ = ⟨q, R_(n−m) k⟩ 对每个 m 都成立 —— 它是分数本身的性质,而不是一组要靠模型去找的基。
每一对都按自己的速率转,所以一个位置就是同时读整把梯子。滑动位置,看这些表盘绕圈:
在上,第 0 对已经转了 32.0 弧度 —— 整整五圈 —— 而第 7 对才走了 0.569。快的表盘分辨相邻,慢的表盘让相隔几百个 token 的词仍然可区分,而这个头一次把八个都读了。
沿整把梯子求和,就得到一个随距离衰减的分数。探一个间隔,再换它底下的底数:
对内容相同的 query 和 key,这就是 §04 那同一个余弦和:RoPE 的衰减和正弦点积是同一个东西。间隔 64 的分数在底数 10000 时是 0.54,在时是 0.67。
而且几乎不要钱:预先算好正弦,Llama-3-8B 的一层花 3 × (4096 + 1024) = 15360 次,对上四个投影矩阵的 8390 万次 —— 占 0.018%,零参数,V 完全不动。
比训练时更长
没有表就没有硬性上限 —— 但不等于模型见过你即将递给它的角度。
一个训练到 2048 个 token 的模型,只给每一对角度落在2048·θᵢ 以内的组合打过分。问它位置 8192 什么都不会报错 —— 每一对只是转到了比训练里任何情况都远四倍的地方。
这就是整个长上下文问题,画出来的样子。拖动缩放,直到被问到的那些角度重新落回模型训练过的那些上:
在上,第 4 对读出的是 204.8 弧度而不是 819.2,正好是训练触到的天花板。这就是位置插值,它用 1000 步微调把 LLaMA-7B 从 2048 撑到了 32768 个 token。账单是分辨率:相邻两个 token 之间的角度差只剩原来的四分之一。
机制本身在一个头里就是两行,而且从不碰 V:
q, k, v = x @ Wq, x @ Wk, x @ Wv q, k = rope(q, pos), rope(k, pos) # v is not a = softmax(q @ k.mT / d_head**0.5) @ v
要知道的那个坑。世面上有两种配对方式:论文把槽位 2i 和 2i+1 配成一对,而所有 HuggingFace 的 LlamaAttention 把 i 和 i + d/2 配成一对。两者互为一个置换,所以各自单独看都是对的。
# the paper pairs (2i, 2i+1) q = interleave_rotate(q, cos, sin) # HuggingFace pairs (i, i + d/2) q = q * cos + rotate_half(q) * sin
拿头里的一个槽位在两种约定下各追一遍。论文那道括号和HuggingFace 那道在槽位 0 上是一致的,之后就不再一致:
到,一种约定给它 θ2,另一种给 θ5 —— 慢了三十倍的频率。拿错的那种去跑一份权重,什么都不会报错:损失照样有限,文本照样通顺,质量悄悄掉下去。
这一切都不花参数。滑动学习表必须覆盖的上下文长度,跟另外两种比一比:
GPT-2 的表为 1024 个位置花掉 786432 个参数;到 就是 25165824 个,在位置 32769 上仍然毫无用处。正弦式和 RoPE 全程贴着零那条轴。
所以:正弦式超出训练长度后无声失效,学习表在 max_len 处大声失效,RoPE 也无声 —— 但只有它在失效上装了旋钮。