自己注意 入門

自己注意を、見えるほど小さく描く。6 トークンの文、ベクトルは 2 次元、すべてのクエリ・キー・バリューが 1 枚の平面に載る。スコア行列、√d の除数、softmax、因果マスク、加重和 —— 数値はすべて隣の図が計算している。マスクを誤って適用する図の数値も同じだ。

01

1 本のベクトルから 3 本へ

トークンは 1 本のベクトルとして到着する。注意の最初の一手は、それをクエリ、キー、バリューに分けることだ。

例文は the cat ate and it ran、そしてこのページはその中の 1 つ、代名詞 it を追いかける。各トークンは実際のヘッドが使う 64 次元ではなく 2 次元のベクトルとして描く。すべての量が平面の上で指させるものになり、議論そのものは幅に依存しないからだ。

3 つの学習済み行列が、この 1 本を 3 本に変える。W_Q はこのトークンのクエリを、W_K はキーを、W_V はバリューを与える。スライダーで 1 本ずつ装着してみてほしい:

入力ベクトルのみ

同じ入力から出た 3 本が、3 つの別々の向きへ伸びることに注目してほしい。クエリはこのトークンが探しているもの、キーは外に向けて公告するもの、バリューは実際に選ばれたときに渡すものだ。名前は情報検索からの借用で、この類比はページの最後まで持ちこたえる。

3 つの行列はこの層では固定なので、トークンを動かせばクエリもキーもバリューも一緒に動く。灰色の入力ベクトルを平面上の好きな場所へ:

入力ベクトルは [1.10, 0.45]. 平面上をドラッグ;方向キーで 0.05 ずつ動き、Home で初期位置に戻る。
入力ベクトルは [1.10, 0.45]

入力を上へ押すとクエリが大きく振れるのに、バリューはほとんど向きを変えない。これが 3 行列で買えるものだ。ある語が何を探すか、何を公告するか、何を渡すかはその語の 3 つの別々の関数であり、1 本のベクトルでは 3 役を同じ役にしてしまう。

他のトークンも同じ平面に自分のキーを持っている。 6 つを順に送って、各クエリがどのキーに最も長い影を落とすか見てほしい:

「it」のクエリ

では勝者は cat のキーで、1.72 対 it 自身のキーの 1.14 だ。it が cat を指すとはだれも教えていない。勾配降下が W_Q と W_K を、代名詞のクエリが名詞のキーへ傾くまで形づくった。それが損失を下げたからだ。

2 つの写像が 1 つだったら、クエリは自分自身のキーになり、尋ねることと尋ねられることが同じになる。W_Q を W_K に混ぜて、2 つの向きが重なるのを見てほしい:

W_Q を W_K に 0% 混合

触る前は cat が ate に 0.51、ate が cat に −0.32。動詞が主語を求める度合いは、その逆よりずっと強い。ところで両方 1.16 になり、関係は類似度に成り下がる。言語は対称ではない。2 つの写像を分けて買っているのは、この向きだ。

02

スコアとは、そろい具合

ベクトル 2 本が入って、数が 1 つ出る。その数は、クエリがキーに沿ってどれだけ伸びているかだ。

算術としては内積、対応する成分を掛けて足すだけ。幾何としては影だ。クエリの先端からキーの直線へ垂線を下ろすと、スコアはその影の長さ × キーの長さになる。クエリをドラッグして数を見てほしい:

クエリ · キー = 1.72. 平面上をドラッグ;方向キーで 0.05 ずつ動き、Home で初期位置に戻る。
クエリ · キー = 1.72

符号がどこで反転するかに注目してほしい。クエリがキーと同じ側に傾いている間は正、直角でちょうど 0、反対側へ倒れると負になる。負のスコアは誤りではなく、そのトークンが反対票を投じているということだ。

1 つのクエリは 6 つのキー全部に同時に当てられ、平面は 6 個の数の並びになる。クエリを文の中で動かして、柱からその行を読み取ってほしい:

「it」の行

6 つのスコアは生の内積なので、上にも下にも際限がない。は自分のキーで 3.02 まで届き、は cat で −0.45 まで落ちる。ここまでのどこにも「大きい」の基準はない。

そして生の内積の罠もそこにある。2 つの長さの積なので、キーは有用な向きを指さなくても、長いというだけで勝てる。and のキーを伸ばして、上がっていくのを見てほしい:

「and」のキーを 1.0 倍に伸ばす

で、機能語の and が 1 度も向きを変えないまま cat を追い抜く。実際の Transformer は射影の前の LayerNorm でこれを抑え、最近のモデルには q と k 自体を正規化するもの(QK-norm)もある。長すぎるキーが 1 本あるだけで行列の全行を奪えるからだ。

03

すべてのクエリがすべてのキーに当たる

クエリが行、キーが列。この正方形が Q · Kᵀ のすべてだ。

検索も候補の絞り込みもない。すべてのトークンのクエリが、自分自身を含むすべてのトークンのキーと内積を取り、結果は文の長さを一辺とする正方形に並ぶ。1 行ずつ埋めてみてほしい:

6 行のうち 0 行

各正方形の一辺がスコアの大きさなので、数字を 1 つも読まずに絵が読める。 6 トークンで 。 1000 トークンの文脈なら 100 万回、しかもヘッドごと層ごとにだ。みんなが文句を言う二次のコストはこれだが、GPT-2 small では 4 つの射影のほうが2·d_model、つまり 1,536 トークンを越えるまで高くつく。

1 行は 1 つのクエリと文全体の対戦表 —— §02 で柱から読んだ 6 個の数を、他の全員の分と積み重ねたものだ。行を下へたどってほしい:

「it」の行

は最大の正方形が cat の下に 1.22 で立ち、はほぼ平らだ。限定詞にはとくに探すものがなく、ほぼ平らな行とは注意がとくに好みはないと言っている姿にほかならない。

この正方形を類似度表として読みたくなるが、そうではない。対角線で折り返して、絵がどう変わるか見てほしい:

Q · Kᵀ

W_Qᵀ W_K が対称でないので、cat → ate の 0.51と ate → cat の −0.32 は、同じ組についての別々の数になる。動詞は主語へ手を伸ばし、主語は動詞に反対票を投じる。類似度行列にはこの区別がつかない。この非対称こそ、写像が 1 つでなく 2 つある理由だ。

04

スコアから混合へ

上下限のない 6 個の数が入り、和が 1 になる 6 個の正の重みが出てくる。

softmax は各スコアを指数に載せ、合計で割る。指数はスコアの符号に関わらずすべての項を正にし、合計で割ることで行の和が 1 になる。ティールの階段は重みを左から右へ足し上げ、上端にぴたりと着地しなければならない:

「it」の混合

混合がどれだけ「柔らかい」かに注目してほしい。で最大の重みは cat の 0.30、最小は the の 0.11 —— 3 倍差であって、参照表の引き当てではない。注意はほとんど 1 つを選ばない。全部に少し、どれかに多めに注ぐ。

6 つは独立した 6 つの決定ではない。分母を共有しているので、どれか 1 つのスコアを上げれば、他のすべての重みが必ず下がる。cat のスコアを行に沿ってドラッグし、残り 5 本が譲るのを見てほしい:

「cat」のスコアは 1.22. 左右にドラッグ;方向キーで 0.1 ずつ動き、Home で初期スコアに戻る。
「cat」のスコアは 1.22

右の読み値を見てほしい。1 つのスコアに何をしても、和は 1.000 のままだ。これが機構全体の不変条件で、そのまま表明として書く価値がある —— 各行で all(w > 0) かつ sum(w) == 1。これが出力を任意の線形結合ではなく加重平均にしている。

その表明の下限は厳密だ。スコアはいくらでも下げられるが、重みはゼロへ近づくだけで決して到達しない。押し下げて、指数を読んでほしい:

「cat」のスコアは 1.20

で重みは 1.2e-14 —— 棒では見えず、算術ではまだ正だ。マスクが「とても小さい負の数」ではなく −∞ を使わなければならない理由はここにあり、softmax を exp(s − max(s)) と実装して安全な理由も同じだ。平行移動は比の中で打ち消えるので、安定版は近似ではなく同じ関数そのものになる。

05

式に √d がいる理由

この除数は帳尻合わせの係数ではない。割られている量のばらつきそのものだ。

d 次元ベクトル 2 本の内積は、d 個の積の和だ。成分が独立で平均 0、分散 1 なら、各積の分散は 1、和の分散は d —— つまりスコアのばらつきは √d で増える。幅を動かして曲線から読み取ってほしい:

幅 d = 1

幅の軸は対数 —— 目盛 1 つで 4 倍 —— だが、ばらつきは目盛ごとに倍になるので、曲線はここでも上へ反る。1 のところのアンバーの線は、 softmax がまともに動くために必要なばらつきだ。 ではばらつきは 8.00 —— 8 倍広い。

8 倍広いことは、8 倍うるさいことではない。softmax は指数関数なので、すべてのスコアを 8 倍することは、最大のものをライバルに対して 8 乗することに等しい。幅を上げて、混合が潰れるのを見てほしい:

幅 d = 1、除数なし

の時点で 1 つの重みが 0.991 を占め、64 では全部を取る。これは適当な選択より悪い。softmax がある重みを通して返す勾配は w(1 − w) に比例するので、1.000 に張り付いた行はほとんど何も返さず、どのトークンを選ぶべきだったのかを学ぶのをやめてしまう。

すべてのスコアを √d で割れば、幅がいくつでもばらつきはちょうど 1 に戻る。除数を入れたまま同じスライダーを動かしてほしい:

幅 d = 1、√d で除算

混合は動かない。1 から 256 までどの幅でも 0.572 のままだ。育っていた尺度が割り落とされたからだ。代わりに d で割れば行き過ぎで、ばらつきは 1/√d まで潰れ、どの行も一様な平均へ平らになる。この玩具ヘッドは d = 2 なので除数は 1.414、 GPT-2 small はヘッドごと d = 64 で、ちょうど 8 で割る。

06

マスクと、その定番の壊し方

言語モデルは次のトークンを当てるように訓練される。そのトークンを先に読むことを、注意は止めない。

§03 の正方形は位置 0 が位置 5 に注意することを許す。訓練時にはそれが、モデルが当てろと言われているトークンだ。だからデコーダは softmax の前に対角線より上をすべて消す。行を 1 つずつ取り除き、未来が消えるのを見てほしい:

6 行のうち 0 行をマスク

破線の正方形は、計算されてから捨てられたスコアだ ——。融合していない実装が 21 個を残すために内積を 36 回やる理由がこれだ。マスクは固定でパラメータを持たず、機構全体で唯一、位置というものを知っている部分でもある。

生き残った分は再正規化される。行の和は 1 でなければならないので、残ったトークンが混合をそのまま分け合う。尋ねる位置を文に沿って下ろしてほしい:

位置 0 が使えるのは 6 個中 1 個

では混合は the の 1.00 —— 最初のトークンには候補が 1 つしかないので、クエリが何と言おうと出力は自分のバリューそのものだ。 では重みが 5 つに広がり、cat が 0.34 を取る。

ここで例の間違いだ。0/1 マスクは行列であり、掛けるのがいちばん自然な適用法に見える。コントロールを + (−∞) から × 0 へ切り替えてほしい:

きっかりゼロ

マスクされたスコアが −∞ ではなく 0 になり、exp(0) = 1 なので、ran は本来いないはずの混合で 0.092 を持ち続ける —— モデルが答えを読んでいる。例外は飛ばず、形も正しく、訓練損失はいつもより速く下がる。そこが尻尾だ。正しい行は s = s.masked_fill(m, float("-inf")) で、重み行列のマスク位置がちょうどゼロであることを表明として書く価値がある。

07

その和と、その請求書

重みは答えではなかった。バリューベクトルを注ぎ合わせるときの配合比だ。

最後の一手は 1 行だ。各バリューベクトルにその重みを掛けて、6 本を足す。どの重みも 1 の一部なので、各項はそのバリューの向きへの短い一歩になる。項を 1 つずつ足してほしい:

6 項のうち 0 項

この鎖が一度も折り返さないことに注目してほしい。折り返させる負の重みが存在しないからだ。 6 歩のあと、が it の新しい表現になる。[0.82, 0.81] —— 猫について少し聞かされたトークンだ。

これは §04 の表明の幾何的な姿でもある。正の重みの和が 1 ということは、出力が凸結合だということで、だからバリューが張る多角形の内側に落ちる。cat のスコアを好きなだけ引っぱって、外へ出せるか試してほしい:

「cat」のスコアは 1.22. 左右にドラッグ;方向キーで 0.1 ずつ動き、Home で初期スコアに戻る。
「cat」のスコアは 1.22

出せない。と出力は [1.21, 0.93] で止まる。cat 自身の値に 100 分の 1 だけ届かない。 §04 の厳密に正な下限こそが、角に近づけても決して到達できない理由だ。 1 つの注意ヘッドは、文の中にすでにあるものの平均しか返せない。後ろに非線形の MLP が付く理由はこれであり、 MLP のない注意層を積んでも 1 つの線形写像に近いものへ潰れる理由も同じだ。

次は、この機構にできないことだ。和は集合の上で取られるので、トークンを並べ替えても重みが並べ替わるだけで、他は何も変わらない。読み順をかき混ぜて、出力を見張っていてほしい:

読み順 the cat ate and it ran

棒は並べ替わり、[0.82, 0.81] は動かない。自己注意は置換同変だ。位置情報がなければ the cat ate and it ran とその並べ替えを区別できず、しかも黙って失敗する —— モデルは訓練され、損失は下がり、語順は表現に一度も入ってこない。位置エンコーディングはこの対称性を壊すために存在し、前節の因果マスクは、このページで順序を知っているもう 1 つのものだ。

請求書は、順序対 1 つにつき内積 1 回、それを 2 度 ——Q · Kᵀ と混合だ。つまり 2n²d 回の積和、そして行列を実体化するなら n² 個分のメモリが、トークン自身の占める n · d に対して要る。文脈長を動かしてほしい:

n = 1,024 トークン、2.0 MiB

では、 1 ヘッドのスコア行列は fp16 で 2.0 MiB —— 12 層 12 ヘッドで 288 MiB。では 32.0 GiB、材料のトークンは 192.0 MiB。171 倍の開きだ。 FlashAttention はこのバイト列を一度も書かないために存在する。

08

4 行と、それが描くもの

このページのすべての考えは、この 4 行のどれかだ。

s = q @ k.transpose(-2, -1) / d_k**0.5
s = s.masked_fill(mask, float("-inf"))
w = s.softmax(dim=-1)
out = w @ v

文全体に対して回すと、結果は 1 枚の絵になる。6 行の重み、各行は 1 つのクエリのキー上の混合で、未来は切り落とされている。マスクを外して、右上の三角が戻るのを見てほしい:

各行は自分と過去だけから取る

これがどの論文にも載る注意マップで、いまや 1 マスずつ読める。因果マスクの下で it は cat から混合の 0.34 を取り、マスクなしでは 0.30 だ。壊し方は 4 つ、どれも例外を投げない。マスクを足さずに掛ける、位置情報がない、√d の除数を落とす、キーを正規化しない。