多頭注意 入門
自己注意は各クエリに 1 行の softmax —— 1 つの混合を与える。実際の文がトークンに求めるものは、1 つの混合で正直に答えられる以上だ。多頭注意の仕掛けを、自分でドラッグできる図で追っていく:同じ幅を h 個に切り、h 回の注意計算を独立に走らせ、実際のヘッドのように専門化させ、また 1 つに戻す —— かかる費用は元の幅広ヘッド 1 つ分と同じだ。
1 行、1 つの混合
自己注意は各クエリに 1 行の softmax だけを与えていた。このページは、その 1 行では足りなくなる瞬間の話だ。
クエリの出力はバリューベクトルの加重平均であり、重みは 1 回の softmax から来る。これは実装の細部ではない —— 1 つのヘッドが生み出せるものの形そのものだ。1 つの混合、それが 6 個のキーにどう配られようと。
引き続き it を見る。直前に来たトークンだけに頼るヘッドを 1 つ用意し、クエリを文中の好きな位置へドラッグしてほしい:
このヘッドは文法について何の考えも持たないことに注目してほしい —— it が代名詞だとは知らず、and が 1 つ前の位置にいることしか知らない。これは本物で有用な信号だ。だが it が必要としているのはこれだけではない。
今度はクエリが文法的に依存する相手に頼るように作った、2 つ目のヘッドだ。クエリと一緒に 2 行が動くのを見てほしい —— では、2 つのヘッドはまったく別のトークンを指す:
it は ran の主語なので、2 つ目のヘッドは前方に届く —— 1 つ目のヘッドと同じ 1 トークン分の距離で、向きだけが逆だ。1 つ目のヘッドの規則は後ろしか見られず、前方は表現しようがない。1 行に、本物で、しかも別々の答えが 2 つ。
1 つのヘッドに第 3 の選択肢はない。混ぜるしかない。混合比を純粋な位置から純粋な文法へ動かし、中点を通るときのピークを読んでほしい:
ピーク重みが中点を通るときに落ち、もう一方の純粋な答えへ向かって再び上がるのを見てほしい —— その途中では、この行はどちらのトークンにも本当には近づいていない。、どちらにも近くない。
このへこみはこの文だけの偶然ではない。混合比そのものに対してプロットすると、このヘッドがどちらかの信号に出せる最高の重みは、両方を最も頑張って両立させようとしているまさにその場所で最も低くなる:
妥協は第 3 の答えではない —— 元の 2 つの答えそれぞれの、より劣化した写しにすぎない。次のページでの解決策は、1 つのヘッドを賢くすることではない。1 つのヘッドに 2 つの仕事を背負わせるのをやめることだ。
同じ幅を、h 個のヘッドに切る
2 つのヘッドとは 2 行の softmax という意味だ。トークンが太くなるという意味ではない。
トークンは相変わらず 1 本のベクトルとして届く。幅は d_model 個の数値 —— 下の図では 8、GPT-2 small では 768 だ。マルチヘッド注意はこの幅に何も足さない。h 個の等しい断片に切るだけで、各断片はd_k = d_model / h 幅になる。
ドラッグして h を変え、同じ幅が違う切られ方をするのを見てほしい —— は切れ目が 1 つもない、1 つのヘッド、ベクトル丸ごとだ:
切れ目を 1 つ増やすたびに各断片が狭くなる一方で、両端は一度も動かないことに注目してほしい。幅 2 のヘッド 4 つは、次元を足し合わせれば幅 8 のヘッド 1 つとまったく同じ予算だ —— 並び方が違うだけだ。
ここで招きやすい誤解がある。各ヘッドがトークン自身の 4 分の 1 を読んでいると想像したくなる。2 つの絵を切り替えて、各ヘッドが出力ではなく入力で実際に何を見ているか確かめてほしい:
各ヘッドの射影はそれぞれ自前の完全な d_model × d_k 行列なので、各ヘッドは 8 個の入力数値すべてを読み、細い結果だけを残す—— トークン自体が刻まれたことは一度もない。
切られているのは W_Q 自身の出力列だ —— そしてW_K、W_V の出力列も、同じやり方で、それぞれ独立に。もう一度 h を変えて、3 つが一緒に切られるのを見てほしい:
クエリ・キー・バリューはそれぞれ自分の d_model × d_model 行列と自分の切り方を持つ —— 3 つの間で共有されているのは数字 h だけだ。3 節では、この細い断片たちが単独になったとき何をするかを見る。
ブレンドなし —— h 回の完全な計算
細い断片はそれぞれ単独で注意機構の全体を実行する。自分のスコア、自分の softmax、自分の加重和。
d_k = 2 のクエリとキーを持つヘッドは、注意の簡易版ではない —— それは注意そのものであり、幅が狭いだけだ。スコアリングも softmax も加重和も、複数のうちの 1 つだからといって何も変わらない。
クエリをドラッグし、位置を追うほうと文法を追うほうが同時に全力で答えるのを見てほしい。自分の行との間には何も挟まっていない:
どちらの行も相手に向かって和らぐことは一度もないことに注目してほしい。1 節の妥協は消えている —— 直ったのではなく、置き換えられたのだ。独立した 2 つの混合は、あの妥協した 1 行と同じコストで手に入る。
これは 2 つで頭打ちにはならない。1 節では試せなかった 2 つのヘッドも足してみる —— クエリが何であろうと常に同じトークンへ引っ張るヘッドと、どちらにも本気を出さない 4 つ目だ:
4 行のうち 3 行には、鋭く名前の付けられる物語がある。4 つ目にはない。ここに色が付いていないのはまさにそのためだ —— 4 節では、実際に訓練されたモデルでこの 4 種が何を意味するのかを見る。
独立性がもう 1 か所で牙をむく。各ヘッドは自分自身の√d_k で割るのであって、モデル全体のではない。と、d_k が縮むにつれて各ヘッドは静かに平らになっていく:
正しい除数は幅をちょうど打ち消すので、同じ強さの信号に対するピーク重みはどの d_k でも一定のままだ。誤った除数はその信号を薄めてしまう —— エラーもクラッシュもなく、ただ softmax が自信を持てる機会を永遠に失うだけだ。
実際のヘッドは何になるのか
1〜3 節で手作りした 2 種類のヘッドは、話を都合よく進めるための作り話ではない。実際に訓練されたヘッドは、勝手に名前の付けられる型へ分化する。
Voita、Talbot、Moiseev、Sennrich、Titov は翻訳モデルを訓練し、そのヘッドが実際に何に注意しているかを調べた。ほとんどのヘッドは 3 つの役割のどれかに収まる —— このページにすでに出てきた 2 つに加えて、3 つ目:クエリが何であろうと、文中で最も希少なトークンへ引っ張る役割だ。
位置を追うヘッドと常に最も希少なトークンへ手を伸ばすヘッドを並べて、クエリを文全体でドラッグしてほしい:
上の行のピークが毎回クエリの 1 歩後ろを歩き、下の行はまったく動かないことに注目してほしい。同じ 6 個のキー、同じコントロール、構造的にまったく異なる 2 つの規則だ。
すべてのクエリ位置についてプロットすると、この違いは形になる。位置ヘッドはまっすぐな対角線を描き、希少語ヘッドは cat に固定された水平線を描く:
統語ヘッドの線はそのどちらでもない —— 各位置で実際の依存関係がどこかに応じて跳び、it では 1 つ先へ、ate では 1 つ後ろへ跳ぶ。単独でステップさせ、この目標がどうしても 1 つの規則に収まらないのを見てほしい:
このヘッドが答えているのは「これは何にかかるか」であって「何個前か」ではないので、その目標は文法が動く通りに動く。これこそ、1 節の 1 行が位置だけでは真似できなかった信号だ。
見返りはここにある。これほどきれいな物語を持つヘッドこそ、モデルが失う余裕のないヘッドだ。Voita らの翻訳モデルの 48 個のヘッドのうちと、専門化したヘッドが生き残る:
生き残ったヘッドの大半は位置・統語・希少語のいずれかだ —— 3 節に出てきた、物語を持たない 4 つ目の型こそ、真っ先に切られる。専門化したヘッドは自分の幅に見合う働きをする。そうでないヘッドは、ほとんどただ同然で取り除ける。
4 つの出力を、1 つに戻す
各ヘッドはそれぞれ単独で計算を終えた。上の層が必要としているのは、相変わらず 1 本の d_model 幅のベクトルであって、4 本の細いものではない。
素直なやり方がそのまま正解だ。4 つの出力を横に並べる。と、幅はちょうど 2 節で切る前の数に戻る:
最後の空きスロットが埋まり、他は何も動かないことに注目してほしい —— 連結はすでに置かれた数値には触れない。次のヘッドの 2 つの値を末尾に足すだけだ。
これらはプレースホルダーではない。トークン it について、各ヘッドの出力は同じ 6 個のバリューベクトルに対する本物の加重和であり、重みだけがそのヘッド自身のものだ。別のトークンにスクラブしても、8 個の数値はすべて本物だ:
4 つのヘッド、4 つの正直な答え、端から端まで並んでいる —— だがここまでは、それでも 4 つの独立した意見のままだ。ヘッド 1を決めるときにヘッド 3 の数値を読んだ者はまだいない。
そのためにあるのが W_O だ。これももう 1 つのd_model × d_model 行列で、連結全体に一度に作用する。その混合をゼロから上げていくと、あるスロットが全ヘッドの答えの一部を帯び始めるのを見てほしい:
W_O はブロック対角ではないので、位置ヘッド自身のスロットは最終的に 4 つのブレンドになる —— 機構全体の中で唯一、統語的な発見と位置的な発見が同じ出力の数値に影響を与えることを許されている場所だ。
ここまでで 4 つの行列が仕事を終えた。W_Q、W_K、W_V がトークンをヘッドへ分け、W_O がヘッドを元に戻す。4 つとも同じ形をしている:
この共通の形は偶然ではない —— それが次節の議論そのものだ。4 つの行列が 1 つの幅を守り、その幅がいくつのヘッドに分けられていようと関係ない。
h 個のヘッドは実際いくらかかるのか
h 個のヘッドは 1 個の h 倍かかると思い込みやすい。実際はまったく同じだ —— パラメータも計算量も。
h がいくつであろうと、W_Q、W_K、W_V、W_O はそれぞれ d_model × d_modelだ —— 5 節ですでにこの 4 つを同じ大きさに描いた。ここで h をドラッグして、この式が出すパラメータ数を見てほしい:
切れ目が増えるたびに d_k は変わるが、その数字には一度も触れない。4 つの行列が 1 つの固定幅を守っている限り、パラメータ数の式にはh がどこにも出てこない —— おおよそ一定なのではなく、厳密に一定だ。
注意の行列積は d でスケールし、幅 d_model / h のヘッドが h 個集まればちょうど d_model になるので、合計は h × d_model / h 回の積和 —— h は打ち消し合う。h に沿ってドラッグし、この曲線が下に引かれた基準線から離れようとしないのを見てほしい:
と細いヘッド 4 つは、どの文脈長でも同じ曲線を描く —— これは丸めた例示ではなく、GPT-2 small 自身の数値だ。4 つの射影行列はさらに約 1.5 倍かかるが、これもやはり動かない —— この節のコストは、幅がどう切られているかに一切依存しない。
だからといって h がタダというわけではない。上げていけばd_k は縮み続ける —— 同じスライダーを、今度はコストではなく、1 つのヘッドがまだ何を表現できるかとして読んでほしい:
d_k = 1 では、クエリもキーも 1 個の数値になり、同符号の 2 数は常に「同じ向き」を指す —— 区別できる方向はもう 2 つしかない。エラーは出ない。FLOPs もパラメータ数も完全に一定のままだ。ヘッドはただ静かに語れなくなっていく。だから実モデルは d_k を 64〜128 に保ち、h を d_model に寄り添わせて増やす。
7 行と、それが買うもの
このページのすべての考えは、この 7 行のどれかだ。
q = (x @ Wq).view(n, h, dk).transpose(0, 1) k = (x @ Wk).view(n, h, dk).transpose(0, 1) v = (x @ Wv).view(n, h, dk).transpose(0, 1) s = q @ k.transpose(-2, -1) / dk**0.5 w = s.softmax(dim=-1) out = (w @ v).transpose(0, 1).reshape(n, h * dk) out = out @ Wo
view と transpose が分割のすべてだ —— x 自体は何も変わらず、変わるのは 3 つの射影自身の出力の読み方だけ。d_model**0.5 ではなく dk**0.5 なのが 3 節の落とし穴だった。最後の reshape と @ Wo が 5 節の連結と混合だ。h はちょうど 4 回現れ、コストはこの 7 行のどこにも現れない —— 同じ 7 行が、書き換えなしで、Transformer-base の から GPT-3 の まで、実在するどの幅でも動く:
d_k の動きは h よりずっと小さい:3 世代で最も動かなかった比率は、モデルの幅ではなく、ヘッド 1 つあたりの表現の余白だ。幅広ヘッドを複数に割ったのは、計算を増やすためではなく、1 トークンに複数の答えを持たせるためだった。それでもヘッド同士は、どのトークンが先に来たかを伝え合えない。それが次のページの出発点だ。