位置エンコーディング 入門
注意は系列を集合として読む:トークンをシャッフルしても出力はそのまま返ってくる。手で動かせる図で 3 つを示す —— 良いエンコーディングは有界であること、正弦の周波数表は内積を間隔だけの関数にすること、そして query と key を回せばそれがスコアの不変式になること。
注意は何がどこにあるか見えていない
トークンをシャッフルしても、出力は並び替わるだけで中身は同じまま返ってくる。修理は 1 回の足し算 —— そしてその足し算はタダではない。
スコアに何が入っているかを見てほしい。Q[i]·K[j] が読むのは2 つのトークンの中身だけで、softmax(QKᵀ/√d)·V のどこにも添字は現れない。注意は置換同変である:入力の行を並べ替えると出力の行も同じ順で並べ替わるだけだ。
3 つのトークンの 6 通りの並べ方を、各トークンの注意出力を下に描いて並べた。順序をドラッグして、出力がワイヤに沿って自分の語と一緒に動くのを見てほしい:
読み出しが 0.00 から離れないことに注目。「犬が人を噛む」と「人が犬を噛む」は、同じ 3 本のベクトルを並べ方だけ変えて注意に渡している —— 言語にとってそれは誤答だ。上のどの層も、どちらが動作主だったかを復元できない。
そこで、注意が見る前にベクトルへ位置を入れておく。各トークンの埋め込みに、どこにいるかだけで決まるベクトルを足す —— 位置を動かすと、真ん中の行が変わり、上の行は微動だにしない:
注意が読むのは一番下の行だけだ。‖PE(p)‖ がスライダーをどこへ動かしても 4.00 のままなのを見てほしい:各ペアは単位円上の 1 点なので、16 ペアなら長さは必ず 4 になる。どの位置でも有界であることは、エンコーディングが最初に満たすべき条件だ。
しかし 1 本のベクトルが 2 つの事実を運ぶようになり、両者は干渉する。位置の重みを上げて、2 つの類似度を比べてほしい ——同じ語が 5 離れている場合と、違う語が同じ枠にいる場合:
2 本の曲線はで交差する。そこを越えると、違う語が同じ枠にいる場合のほうが同じ語が 5 離れている場合より似てしまい、上の層は区別を失った和をほどく羽目になる。2017 年の論文が足す重みは 1 で、交差点のかなり左側にある。
条件は、有界で、位置ごとに異なり、内容がこの和を生き延びる程度に静かで、できれば距離も語れること。次節では最初に思いつく 2 つの方式を試す。
数えるのは見た目より難しい
誰でも思いつく 2 つのエンコーディングと、それぞれが壊れる具体的な壊れ方。生き残るのは、複数のスケールで同時に数える方式だ。
いちばん単純なエンコーディングは添字そのものだ。位置 0 には 0、位置 1 には 1、以下いくらでも続く。訓練も保存も不要、最大長もない —— そして最も早く壊れる。
足される相手のベクトルの長さはおよそ 1 だ。マーカーを坂に沿って上げて、埋め込みが暮らしている帯をどれだけ早く飛び出すか見てほしい:
では、エンコーディングは足される相手の 63 倍になっている。この和を読む層には位置しか見えず、トークン埋め込みへ届く勾配は押し流される。有界であることは飾りではない。
素直な修正は系列長で割ることで、これは確かに値を抑える。2 本目の系列を伸ばして、枠 3 がそれぞれいくつになるか見てほしい:
枠 3 は8 トークンの系列では 0.43、64 トークンの系列では 0.05 になる。エンコーディングは位置の性質であることをやめ、バッチの性質になってしまった —— 「3 トークン前」に対応する固定の表現が、もうどこにも存在しない。
欲しいのは、有界で、絶対的で、なお距離を語るものだ。2 進法はその 3 つを既に満たしている:位置をドラッグして、各ビットが自分のペースで反転するのを見てほしい:
0 番ビットは 2 トークンごと、3 番ビットは 16 トークンごとに反転する。上から下へ読むと速いビットが隣を、遅いビットが領域全体を分け、値は {0, 1} のままだ。
この走行距離計の唯一の欠点は段でできていることだ。01111 から 10000 へ進むと 5 ビットが一度に反転し、勾配が辿れる方向がない。角を丸めれば正弦エンコーディングになる。
角を丸めた走行距離計
等比に並んだ周波数のはしごに載せた正弦と余弦。パラメータはゼロ、どのスケールにも波長がある。
2017 年の式は、位置 p の枠 2i を sin(p·θᵢ)、枠 2i+1 を cos(p·θᵢ) と定める。ここで θᵢ = base^(−2i/d)。各ペアが自分の周波数を持ち、等比で下がる。
これが角を丸めた走行距離計だ。ペアを 1 つ選び、読み出しからその周波数を読んでほしい —— 上の行は数トークンで一周し、下の行はほとんど曲がらない:
ペア 0 は 1 トークンあたり 1.000 ラジアン回るので、6.3 トークンで一周する。 は 353 かかる。はしごを等比にしたのは意図的だ:どの段も同じ 2 枠しか使わないのに、前の段より一桁粗いスケールを買える。
1 つの位置のエンコーディングは、それら全体を縦に切った断面である。切り口を動かして、8 個の点が高さを見つけるのを見てほしい:
この 8 個の正弦値と、その 8 個の余弦の相棒が、PE(p) の 16 個の数だ。同じ列を共有する位置は 2 つとなく、隣り合う列は速い行でしか違わない —— ビット列が持っていた距離感が、そのまま連続になっている。
はしごがどこまで届くかは底が決める。ペアを 1 つ選んでから、その下の底を変えてほしい:
論文の底 10000 では、このはしごはペア 0 の 6.3 トークンから の 35333 まで伸びる。 16 段のうち 5 段は訓練した範囲より上にあり、そのなかで一周しない。底を上げればはしご全体が伸びる。これが Llama 3 が引いたレバーだ。
こうしてエンコーディングは有界で、位置ごとに異なり、多スケールで、しかも無料になった。残るのは、この先すべてが依存する性質だ —— 2 つの位置の距離について何か言えているのか。
欲しいのは間隔であって場所ではない
各ペアは円周上の 1 点なので、前へ進むことは回ることだ。正弦の周波数表が単なるハッシュ以上のものになるのは、この一点による。
ペアを 1 つだけ取り出す。(sin p·θᵢ, cos p·θᵢ) は角度 p·θᵢ にある単位円上の点であり、p から p+k へ進むことは k·θᵢ の回転にあたる —— そしてその回転量は p にまったく依存しない。
位置を動かすと弧が伸び、前へ進む一歩はぴったり同じ大きさを保つ:
位置がどこで終わってもその一歩の回転量は同じなので、PE(p+k) は PE(p) の線形関数であり、その行列は k だけで決まる。重み行列 1 つで「3 トークン前を見る」をすべての位置について同時に実装できる。
これを 16 ペア全部で同時に行うと、何かが潰れる。波の重なりに縦の切り口を 2 本入れ、それが選ぶ2 本の列どうしの内積を読んでほしい:
2 本の切り口を一緒に動かしても数は動かない。間隔 8 なら、切り口の組がどこにあっても読み出しは 0.66 で、 1.00 になるのは 2 本が重なったときだけだ。
同じ事実を曲線にしたものだ。最初の位置を固定し、それを他のすべての位置と突き合わせる ——曲線全体が形を保ったまま一緒に動く:
sin a sin b + cos a cos b はちょうど cos(a−b) なので、PE(p)·PE(q) は Σᵢ cos((p−q)·θᵢ) になる —— 間隔だけの関数で、それ以外には依存しない。8 トークン先の値は p が 0 でも でも 0.66 のままだ。
ここが要約で飛ばされる部分だ。注意が採点するのは PE 対 PE ではなく、(x+PE)W_Q 対 (x+PE)W_K である。射影を恒等写像から押し出してほしい:
静止状態では 8 本の曲線は厳密な曲線の上にぴたりと重なり、射影が動いた途端に広がる:混合 1.00 では間隔 8 でのばらつきが 0.48、値域のほぼ 4 分の 1 に達する。平行移動不変性はモデルが使える基底であって、必ず手に入る保証ではなかった。
さらに (x+PE)W_Q · (x+PE)W_K の展開は 4 項で、位置対位置は 1 項だけだ。エンコーディングの約束とスコアの実際との隙間こそ、RoPE が入り込む口である。
位置ごとに 1 行
BERT と GPT-2 は算術を飛ばした。位置ごとに 1 行のパラメータを確保し、中身は最適化器に決めさせる。
nn.Embedding(max_len, d_model) を添字で引き、正弦ベクトルとまったく同じように足す。トークン表と同じ経路で語彙が違うだけ —— 埋め込みを書いた人は既に書いている。
行の中身に制約はない。行を 1 つ選び、滑らかさを動かしてほしい —— どちらの端も、勾配降下が残しうる状態だ:
では各行は独立な抽出であり、それこそがパラメータ化の保証する内容だ —— つまり何も保証しない。訓練後の表はたいていこれより滑らかになるが、モデルがそれを要求したわけではない。構造はデータの副産物であって、方式の性質ではない。
それが隣り合う行の姿を決める。学習された類似度を、同じ間隔での厳密な正弦の類似度と並べてほしい:
遅いペアがほとんど動いていないので、正弦の曲線は間隔 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 の次元のペアを 2 次元とみなし、§03 の周波数表に従って p·θᵢ だけ回す。
円上の 2 本のベクトルの内積は、両者の角度だけで決まる。query の位置とkey の位置を動かして、スコアを見てほしい:
両方を同じだけ動かしてもスコアは動かず、片方だけ動かすと間隔だけに応じて変わる。これが不変式だ:⟨R_m q, R_n k⟩ = ⟨q, R_(n−m) k⟩ はどの m でも成り立つ —— モデルが見つけるべき基底ではなく、スコアそのものの性質である。
各ペアは自分の速さで回るので、1 つの位置ははしご全体を一度に読むことになる。位置を動かして、ダイヤルが巻かれるのを見てほしい:
ではペア 0 が 32.0 ラジアン —— 5 周 —— 回っているのに対し、ペア 7 は 0.569 しか進んでいない。速いダイヤルが隣を分け、遅いダイヤルが数百トークン離れた語を区別可能に保ち、ヘッドは 8 つを同時に読む。
はしご全体で和をとると、距離とともに減衰するスコアが得られる。間隔を調べ、その下の底を変えてほしい:
同じ内容の query と key では、これは §04 と同じ余弦の和だ:RoPE の減衰と正弦の内積は同じ対象。底 10000 では間隔 64 のスコアが 0.54、 では 0.67 になる。これこそ Llama 3 が引いたレバーだ。
しかもほぼ無料だ。正弦を事前計算すれば Llama-3-8B の 1 層は 3 × (4096 + 1024) = 15360 フロップ、4 つの射影行列の 8390 万に対し 0.018% だ。パラメータはゼロ、V も回さない。
訓練したより長く
テーブルがなければ上限もない —— だがモデルがこれから渡す角度を見たことがある、という意味でもない。
2048 トークンまで訓練したモデルは、角度が2048·θᵢ の内側にある組にしか点を付けていない。位置 8192 でもエラーは出ず、各ペアが訓練時の 4 倍先まで回って届く。
これが長文脈問題のすべてを絵にしたものだ。スケールをドラッグして、要求されている角度を訓練された角度の上に戻してほしい:
ではペア 4 が 819.2 ではなく 204.8 ラジアンを示す。訓練が到達した天井とちょうど同じだ。これが Position Interpolation で、1000 ステップの微調整で LLaMA-7B を 2048 → 32768 トークンへ広げた。請求書は解像度だ:隣り合うトークンの角度差は 4 分の 1 になる。
仕組み自体はヘッドの内側の 2 行で、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
知っておくべき落とし穴。ペアの組み方は 2 通り出回っている。論文は枠 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
ヘッドの枠を 1 つ選び、両方の規約で追いかけてほしい。論文のブラケットとHuggingFace のブラケットは枠 0 では一致し、そこから一致しなくなる:
では、一方の規約は θ2 を、もう一方は θ5 を割り当てる —— 30 倍遅い周波数だ。誤ったほうで重みを動かしても何も例外は出ない:損失は有限のまま、文章は文法的なまま、品質だけが静かに落ちる。
そしてこのどれもパラメータを使わない。学習テーブルが覆わねばならない文脈長を動かし、他の 2 つと比べてほしい:
GPT-2 のテーブルは 1024 位置に 786432 パラメータを使う。 なら 25165824 で、位置 32769 では役に立たない。正弦型と RoPE は掃引のあいだ軸の上、ゼロのままだ。
まとめると:正弦型は訓練長を越えると静かに壊れ、学習テーブルは max_len で大声で壊れ、RoPE も静かに壊れる —— ただし壊れ方につまみが付く唯一の方式だ。