TT Lab
はじめる
学ぶ 学習パス コース

Transformer — アテンションを手で計算する

位置は足さずに回す

TT Labで続きを見る

一言でいうと

位置をベクトルに足す代わりに、ベクトルを位置の分だけ回転させると、クエリとキーの内積が、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つで描いたRoPE。位置3と位置5のベクトルがなす角と、位置4000と位置4002のベクトルがなす角は、どちらも同じ2θです。2つのベクトルを同じ向きにいっしょに回しても、間の角は変わらないので、内積は2つの位置の差だけに依存します

言葉にすると、こうなります。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周を超えていないペアがいくつ残るかを測ります。採点ツールは、書いておいた説明を信用しません。自分で作ったモジュールを実際に呼び出し、毎回異なる次元と位置で関数を直接叩いて、採点ツールが別に計算した値と、許容誤差の範囲内で照合します。