位置は足さずに回す
一言でいうと
位置をベクトルに足す代わりに、ベクトルを位置の分だけ回転させると、クエリとキーの内積が、2つの位置の差だけに依存するようになります。位置3と位置5で測った値と、位置4000と位置4002で測った値が、同じになります。
なぜ必要なのか
前のラボで、サイン/コサインの位置ベクトルを作り、入力に足してみました。それで順序が入ったことは確かです。問題は、どのように入ったかです。
アテンションのスコアは、クエリとキーの内積1つで決まります。位置ベクトルPを足したあとで内積を展開すると、4つの項が出てきます。
(q + P[m]) · (k + P[n])
= q·k + q·P[n] + P[m]·k + P[m]·P[n]
先頭の項は内容同士のスコアで、末尾の項は位置同士のスコアです。問題は、真ん中の2つの項です。q·P[n]にはnだけが、P[m]·kにはmだけが入っています。片方の位置しか持たない項は、差としてまとめられません。そのため、間隔がまったく同じ2であっても、文の前のほうで測った値と、後ろのほうで測った値が異なります。
これがなぜ悪いのでしょうか。言語で重要なのは、たいてい「何語前の単語か」であって、「文書の何文字目か」ではありません。連体修飾節が修飾する名詞はすぐ後ろにあり、代名詞が指すものは数文前にあります。どれも距離で表される関係です。ところが、モデルが見るスコアには、絶対位置が混ざり込んでいます。学習のときに見たことのない位置番号がサービスで出てくると(より長い文章を入れた瞬間に、まさにそうなります)、それらの項がどんな値になるのか、誰にもわかりません。
直そうとする試みは、以前からありました。位置ベクトルを学習させると、表にない位置でそのまま止まります。スコアに距離ごとに異なる定数を足す方法は、項が1つ増えるだけで、内容と位置が混ざる問題は残ります。どちらも、内積そのものには手を付けられませんでした。
そこで、RoFormer論文が投げかけた問いは、これです。内積が最初から相対位置だけを見るようにできないか。スコアを出したあとで直すのではなく、クエリとキーを作る時点で、すでにそうなっているようにします。
足す代わりに回す
答えは、演算を変えることです。足さずに、回します。
偶数次元を2つずつ組にして、平面上の点として見ます。64次元なら、ペアが32個です。ペアごとに角度を決めておき、m番目の位置のベクトルは、各ペアをm掛けるその角度の分だけ回転させます。
theta_i = base ** (-2*i/d) # i 번째 쌍의 회전 속도
(x0, x1) -> (x0*cos(a) - x1*sin(a), x0*sin(a) + x1*cos(a)) # a = m * theta_i
ペア0は、速度1で最も速く回り、後ろへ行くほど指数関数的に遅くなります。底は、論文の定義では10000です。
ここで重要なのは、足したものが何もないという点です。ベクトルに新しい値を混ぜていません。もともとあった値の向きだけを、位置の分だけ回しました。
ペアの中だけで回る点も、注目に値します。ペア0は次元0と次元1だけを混ぜ、ペア1は次元2と次元3だけを混ぜます。次元全体をかき混ぜる大きな行列ではなく、対角線に2×2の小さな回転がずらりと並んだ形です。そのため、掛け算を一度にまるごと行うのではなく、ペアごとにcosとsinを掛けて足すだけで済みます。
なぜ相対位置だけが残るのか
ペア1つだけを取り出すと、1行で終わります。回転行列Rの転置は、逆方向の回転です。
(R(m*theta) q) · (R(n*theta) k) = q · R(m*theta)ᵀ R(n*theta) k
= q · R((n-m)*theta) k
mとnがそれぞれどこにあるのかは消えて、差だけが残ります。ペアごとにこれが別々に起こり、内積はそれらを足したものなので、ベクトル全体でも同じことが成り立ちます。
言葉にすると、こうなります。2つを同じ向きにいっしょに回しても、間の角は変わりません。時計の針2本をまるごと回しても、2本が開いた角度はそのままであるのと同じです。内積は結局、長さと間の角で決まるので、間の角が変わらなければ、スコアも変わりません。
そして、回転は長さを変えません。cosとsinの2乗を足すと1になるからです。足す方式は、もとの信号の上に別の値を載せて大きさを変えますが、回転は載せるものがありません。位置情報を入れながら、内容には手を付けていないことになります。
離れると何が残るのか
ペアごとに速度が違うことが、ここで効いてきます。
間隔deltaだけ離れた2つの位置の間で、i番目のペアはdelta * theta_iだけ開きます。これを1周(2π)で割ると、何周回ったかが出ます。速いペアは、少し離れただけで何周も回ってしまいます。1周を超えたペアは、deltaと「deltaから1周分を引いた距離」を同じ角度で表します。そのペアだけを見ても、両者を見分けられないということです。
そのため、距離が離れると、まだ1周していない遅いペアだけが、その距離を正しく区別します。近い距離は速いペアが細かく分け、遠い距離は遅いペアが大きく分けます。速度を指数関数的に敷いておく理由が、これです。物差しを何本も重ねておくようなもので、目盛りの細かい物差しは短いものだけを測り、目盛りの粗い物差しは長いものを測る、という具合です。
ラボで64次元を使って実際に測ってみると、数字がはっきり出ます。間隔が1のときは、32個のペアがすべて最初の1周の中にありますが、間隔が大きくなるほど、その数が減っていきます。減っていく様子を目で見ると、「コンテキストを延ばす作業」がなぜ角度を調整する話につながるのかが、見えてきます。残った物差しをもっと長くするか、物差しをもっと粗く敷き直す必要があるからです。
現場での姿
第1に、コンテキストを延ばしたら、品質が崩れます。学習で使った長さを超えると、遅いペアでさえ、初めて見る角度に入ります。角度そのものは計算されますが、モデルがその角度で何をすべきかは、学んだことがありません。
第2に、どの層に入れるのかを混同します。位置ベクトルを足す方式は、入力の埋め込みに一度足せば終わりです。回転はそうではありません。アテンションがクエリとキーを使う箇所で適用されます。値(V)は回しません。回すと、内容が位置に引きずられるからです。
第3に、クエリとキーが異なる約束を使います。ペアを隣同士で組む実装と、前半と後半を対にする実装は、どちらも一般的です。どちらでも性質は同じですが、片方の重みをもう片方のコードに入れると、静かに間違います。エラーは出ず、スコアだけがおかしくなります。
第4に、キャッシュに回転済みの値を入れたのに、位置をもう一度数えます。すでに回転させたキーをキャッシュに置き、あとでもう一度回すと、位置が二重に入ります。生成が長くなるほど、ずれていきます。
第5に、底を変えると、別のモデルになります。底は、ペアたちの速度をまるごと決める値です。学習した底と異なる底で推論すると、すべてのペアの角度がずれます。
実務で本当に大切なこと
- 同じ間隔なら同じスコアになることを、テストで固定してください。位置を動かしながら測って、値が同じかを見るのは数行のテストですが、これが壊れると、以下のすべてが意味を失います。
- 回転はクエリとキーだけに行います。値には入れません。
- 長さが保たれていることを確認してください。回したあとでノルムが変わっていたなら、回転ではなく別のことをしたのです。
- 浮動小数点の誤差を覚悟してください。位置が大きくなると角度も大きくなり、誤差が増えます。同じかどうかを見るときは、
==ではなく、許容誤差を設けた比較を使います。
次のラボですること
/root/work/tf-rope/rope.pyを、1ステップずつ育てていきます。numpyもtorchも使いません。このPodのシステムのPythonにはnumpyがなく(/opt/onnx-lab/bin/pythonの中にしかありません)、インターネットもありません。標準ライブラリのmathだけで十分です。
ペアごとに異なる回転速度を作ることから始めて、2次元の回転1つ、ベクトル全体を位置の分だけ回すこと、回したクエリとキーのスコアまでを作ります。そのあと、足す方式を並べて作り、同じ位置で2つの値を比べます。
中心になるのは、ステップ6です。間隔を2に固定しておき、開始位置を3、10、100、4000と動かしながら、2つの方式でスコアを測ります。回転のほうは、4つの位置で同じ値が出て、足すほうは揺らぎます。そのばらつきの幅を、数字で直接見ることになります。
最後に、間隔が離れるときにペアごとに何周回るかを数え、まだ1周を超えていないペアがいくつ残るかを測ります。採点ツールは、書いておいた説明を信用しません。自分で作ったモジュールを実際に呼び出し、毎回異なる次元と位置で関数を直接叩いて、採点ツールが別に計算した値と、許容誤差の範囲内で照合します。