コンテンツにスキップ

3. Attention

猫が魚を食べた。それは新鮮だった

この「それ」が魚を指すことは、人間なら後ろの「新鮮だった」まで読んですぐ分かります。でも、埋め込みの時点では、「それ」のベクトルは周りの文を何も知りません。

この回では、「それ」のベクトルに、前にある「猫」や「魚」の情報を取り込むしくみを見ます。

Attention は、各トークンが「文中のどのトークンから情報をもらうべきか」を計算し、もらう情報を集めてくる仕組みです。集めた情報は、次の回で見るように元のベクトルに 足される ので、元の「それ」の情報が残る通り道を保ちながら、「魚」の情報が加わります。Transformer の心臓部と言える部分です。

この回の入力と出力
  • 入力: ベクトルの列(10 × d)。前の回の出力、またはひとつ前のブロックの出力(実際には LayerNorm で整えてから入れる。4 の回)
  • 出力: 同じ形のベクトルの列(10 × d)。各行は、そのトークンが自分自身と前のトークンから集めてきた情報(後ろは見ない。下の「未来を隠す」の節)
  • その後: 次の回で、元のベクトルに足される
  • 途中で作るもの: Query・Key・Value の 3 種類のベクトル(それぞれ 10 本)と、「どのトークンがどのトークンをどれだけ見るか」を表す 10 × 10 の重みの表
  • 使うもの: Query・Key・Value を作るための行列数値を縦横の表の形に並べたもの。横の並びを行、縦の並びを列という。ベクトルに行列を掛けると、別のベクトルに変換できる。用語集でくわしく →と、出力を整えるための行列(下のマルチヘッドの節で説明。どれも学習で決まるパラメータモデルが学習で調整する数値。埋め込み行列や、Attention・MLP の行列の中身がすべてパラメータ。「70B のモデル」は、パラメータが 700 億個あるという意味。用語集でくわしく →)

Attention では、各トークンのベクトルから 3 種類のベクトルを作ります。それぞれ、学習済みの行列を掛けて作ります(線形変換ベクトルに行列を掛けて、別のベクトルに変えること。各出力は、入力の数値に重みを掛けて足し合わせたものになる。用語集でくわしく →)。

ベクトルに行列を掛けると、出力の各数値は「入力の数値に重みを掛けて足したもの」になります。たとえば入力が [2, 1] で、行列が「1 つ目の出力 = 1×(1 つ目) + 0×(2 つ目)、2 つ目の出力 = 1×(1 つ目) + 2×(2 つ目)」なら、出力は [2, 4] です。同じ入力でも、Query 用・Key 用・Value 用に別々の行列を掛けるので、3 種類の違うベクトルができます。

ベクトル役割たとえ
Query(q)自分がどんな情報を探しているか検索ワード
Key(k)自分がどんな情報を持っているか検索される側の見出し
Value(v)実際に渡す情報の中身見出しの先にある本文

あるトークン(たとえば「それ」)の出力は、次の 3 ステップで計算します。

  1. スコア: 自分の q と、各トークンの k の内積2 つのベクトルの同じ位置の数値どうしを掛けて、全部足した値。それぞれの長さを変えずに比べれば、向きが揃っているほど大きく、逆向きなら負になる。ベクトルが長いほど、値の振れ幅も大きくなる。用語集でくわしく →を取る。内積は、同じ位置の数値どうしを掛けて全部足した値(下の表に途中式があります)。ベクトルの長さが同じなら、向きが揃っているほど大きくなる。「それ」の q と「魚」の k の向きが揃っていれば、「魚」のスコアが高くなる。
  2. 重み: スコアを dk\sqrt{d_k}(dkd_k は q と k の次元数)で割ってから softmax(ソフトマックス)数値の並びを、すべて 0 以上で合計がちょうど 1 になる「確率の並び」に変える関数。元の値が大きいものほど大きな確率になる。用語集でくわしく → にかけ、合計 1 の重みにする。softmax は、各値を指数関数 exe^x(ee は約 2.718 の定数で、その xx 乗)に通してから、その合計で割る計算です(くわしくは 5 の回)。
  3. 出力: 各トークンの v を、その重みで加重平均値ごとに「重み」を掛けてから足し合わせる平均。重みの合計は 1 にする。重みの大きい値ほど結果に強く効く。用語集でくわしく →する。ベクトルの加重平均は、同じ位置の数値ごとに計算する。重みの大きい「魚」の v が多く混ざる。

小さな数値例で、3 ステップを通して追ってみます。2 次元で、各次元を[動物, 食べ物]とし、トークンは「猫」「魚」「それ」の 3 つだけに絞ったおもちゃの例です(v は k と同じ値にします)。「それ」の q は [0, 2](食べ物を探している)と仮に決めます。

キーk(= v)1. スコア(q・k)÷ 2\sqrt{2}(dk=2d_k = 2)2. 重み(softmax)
猫[2, 0]0×2 + 2×0 = 000.045
魚[0, 2]0×0 + 2×2 = 42.830.769
それ[1, 1]0×1 + 2×1 = 21.410.186

重みの計算:e0=1e^0 = 1、e2.83≈16.95e^{2.83} \approx 16.95、e1.41≈4.10e^{1.41} \approx 4.10 で、合計は約 22.05。それぞれを合計で割ると、約 [0.045, 0.769, 0.186] です(小数 3 桁に四捨五入)。

最後に 3. の加重平均をすると、出力は 0.045×[2, 0] + 0.769×[0, 2] + 0.186×[1, 1] = [0.276, 1.724] ≈ [0.28, 1.72]。「食べ物」の成分が大きく、「魚」の情報が多く混ざっています。

これを 10 個のトークン全部について同時に行います。行列でまとめて書くと、有名な次の式になります(Q,K,VQ, K, V は q・k・v を 10 本ずつ縦に並べた行列、K⊤K^\top は K の転置行列の行と列を入れ替えること。A の転置を Aᵀ と書く。用語集でくわしく →、MM は後述のマスク)。

Attention(Q,K,V)=softmax(QK⊤dk+M)V\mathrm{Attention}(Q, K, V) = \mathrm{softmax}\left(\frac{QK^\top}{\sqrt{d_k}} + M\right) V

ここで計算しているのは、1 組の Q・K・V(下のマルチヘッドの節でいう 1 ヘッド)の分です。v の次元も q・k と同じ dkd_k とするので、出力は 10 × dkd_k になります(10 × d に戻す方法は、マルチヘッドの節で見ます)。行列の形は、次の順に変わります。

Kᵀ(d_k × 10)Q猫が魚を食べた。それは新鮮だった×V=出力10 × d_kスコア → 重み(10 × 10)10 × d_kマスクで隠す位置
Q の行と Kᵀ の列が交わるマスが、その 2 つのトークンの内積(スコア)です。行が「見る側」、列が「見られる側」のトークンで、青い行が「それ」の行です。 斜線のマス(右上。自分より後ろのトークン)を隠して行ごとに softmax をかけて重みにし、V を掛けて出力にします。
  • QK⊤QK^\top: QQ(10 × dkd_k)と K⊤K^\top(dkd_k × 10)を掛けて、10 × 10 のスコアの表にする。ii 行 jj 列は「ii 番目の q と jj 番目の k の内積」
  • softmax: 行ごと(クエリのトークンごと)にかけて、10 × 10 の重みの表にする
  • VV を掛ける: 重みの表(10 × 10)に VV(10 × dkd_k)を掛けて、1 ヘッド分の出力(10 × dkd_k)にする。ii 行目は「ii 番目のトークンの重みで v を加重平均したもの」

上の数値例は、この表の「それ」の 1 行ぶん(図の色のついた行)の計算に当たります(ただし数値例は、見る相手を「猫」「魚」「それ」の 3 つに絞っています)。

Attention の計算を 1 ステップずつ見る
ヘッド
クエリにするトークンを選ぶ(このトークンが、どこから情報を集めるか)
「それ」からの注目度(softmax 後の重み)
猫
18.8%
が
1.5%
魚
65.7%
を
1.5%
食べた
6.1%
。
2.5%
それ
3.7%
は
—
新鮮
—
だった
—
「それ」のクエリベクトル q を動かす(何を探しているか)
計算の中身(k の各次元は 動物・食べ物・動作・機能語。q はこの k と内積を取る)
キーk(= v)q·k÷√d重み
猫[3, 0, 0, 0]3.001.500.188
が[0, 0, 0, 2]-2.00-1.000.015
魚[1, 3, 0, 0]5.502.750.657
を[0, 0, 0, 2]-2.00-1.000.015
食べた[0, 0.5, 3, 0]0.750.380.061
。[0, 0, 0, 1]-1.00-0.500.025
それ[0.3, 0.3, 0, 1]-0.25-0.130.037
出力 = Σ 重み × V
[1.23, 2.01, 0.18, 0.12]
(動物 1.2 / 食べ物 2.0 / 動作 0.2 / 機能語 0.1)
全トークンの注目度マップ(行 = クエリ、列 = キー。行の見出しを押すと、そのトークンを選ぶ)
猫が魚を食べた。それは新鮮だった
猫 100.0%が マスク(見えない)魚 マスク(見えない)を マスク(見えない)食べた マスク(見えない)。 マスク(見えない)それ マスク(見えない)は マスク(見えない)新鮮 マスク(見えない)だった マスク(見えない)
猫 81.8%が 18.2%魚 マスク(見えない)を マスク(見えない)食べた マスク(見えない)。 マスク(見えない)それ マスク(見えない)は マスク(見えない)新鮮 マスク(見えない)だった マスク(見えない)
猫 50.0%が 11.1%魚 38.9%を マスク(見えない)食べた マスク(見えない)。 マスク(見えない)それ マスク(見えない)は マスク(見えない)新鮮 マスク(見えない)だった マスク(見えない)
猫 32.3%が 7.2%魚 53.3%を 7.2%食べた マスク(見えない)。 マスク(見えない)それ マスク(見えない)は マスク(見えない)新鮮 マスク(見えない)だった マスク(見えない)
猫 29.4%が 1.9%魚 62.3%を 1.9%食べた 4.5%。 マスク(見えない)それ マスク(見えない)は マスク(見えない)新鮮 マスク(見えない)だった マスク(見えない)
猫 10.5%が 10.5%魚 10.5%を 10.5%食べた 47.3%。 10.5%それ マスク(見えない)は マスク(見えない)新鮮 マスク(見えない)だった マスク(見えない)
猫 18.8%が 1.5%魚 65.7%を 1.5%食べた 6.1%。 2.5%それ 3.7%は マスク(見えない)新鮮 マスク(見えない)だった マスク(見えない)
猫 26.7%が 3.6%魚 44.0%を 3.6%食べた 7.6%。 4.6%それ 6.3%は 3.6%新鮮 マスク(見えない)だった マスク(見えない)
猫 3.1%が 3.1%魚 61.7%を 3.1%食べた 5.1%。 3.1%それ 4.1%は 3.1%新鮮 13.8%だった マスク(見えない)
猫 7.9%が 3.7%魚 45.2%を 3.7%食べた 11.4%。 3.7%それ 5.0%は 3.7%新鮮 11.4%だった 4.2%

デモには 2 つの「ヘッド」があります。ヘッドとは、1 組の Q・K・V を使う Attention のことです(くわしくは下のマルチヘッドの節)。

  • クエリ「それ」の重みを見る: 「魚」に最も強く注目しています。これは「それ」の q が「食べ物らしさ」の方向を向き、「魚」の k も同じ方向を向いているからです。
  • 「それ」の q の「食べ物」スライダーを 0 にし、「動物」を上げる: 注目先が「猫」に移ります。q は「何を探すか」を表している、ということが体感できます。
  • スケールの強さを動かす: このスライダーは、dk\sqrt{d_k} で割ったあと、さらに割る数です(実際の Attention にはない、デモのためのつまみ。5 の回の temperature と同じ働き)。大きくすると重みが平らに(いろいろなところを少しずつ見る)、小さくすると 1 か所に集中します。
  • 「未来を隠す」を外す: 「未来を隠す」は、各トークンが自分より後ろのトークンを見られないようにする仕組みです(理由は次の節)。外すと、「猫」が後ろの「魚」や「食べた」も見られるようになります。

LLM は「次のトークンを予測する」ように学習します。学習のときは、完成した文章を丸ごと入力し、すべての位置で「次のトークン」を同時に当てさせます(そのほうが効率よく学習できる)。このとき、たとえば「猫」の位置から後ろの「が」が見えたら、答えを見ながら解くのと同じです。 そこで、各トークンは 自分より前(と自分自身)しか見られない ようにします。具体的には、未来の位置のスコアを −∞-\infty にしてから softmax にかけ、重みを 0 にします(e−∞=0e^{-\infty} = 0)。式の MM がこれで、10 × 10 の表のうち、見てよい位置(自分と前)は 0、未来の位置は −∞-\infty になっています。スコアの表に足すと、未来の位置だけが −∞-\infty になります。

冒頭の例文でいうと、「それ」の位置からは後ろの「は新鮮だった」が見えず、集められるのは前の「猫」や「魚」の情報までです。その代わり、後ろの「新鮮」の位置からは「それ」も「魚」も見えるので、「新鮮なのは魚」という手がかりは、そこで集められます。

上の数値例では、「それ」の q を「食べ物を探している」向きに仮に決めましたが、この向きは「新鮮」からは来られません。本物のモデルでこう向くとしたら、前のブロックで「魚を食べた」などの前の文脈を取り込んだ結果です(埋め込み直後の最初のブロックでは、「それ」はまだ周りを知りません)。

1 組の Q・K・V では、1 つのトークンにつき 1 通りの重みづけしかできないので、「代名詞の指す先」と「直前のトークン」のように違う種類の関係を同時に集めにくくなります。そこで実際には、別々の行列を持つ Attention(ヘッド)を何個も並列に動かします(数はモデルによります。GPT-2 最小構成では 12 個)。

  • あるヘッドは「代名詞が指すもの」を探す
  • 別のヘッドは「直前のトークン」を見る(デモの「直前を見るヘッド」)
  • また別のヘッドは「文法的な主語」を探す

各ヘッドは、d より小さい次元の q・k・v を使います(v の次元は一般には q・k と別にもできますが、GPT-2 など多くのモデルでは同じです)。よくあるのは、d をヘッド数で割った次元にする方法です。たとえば GPT-2 最小構成では、d = 768 を 12 ヘッドで割って、各ヘッドは 64 次元。

12 個のヘッドの出力(各 10 × 64)を横に並べて(連結して)、10 × 768 に戻します。それにもう 1 つの行列(各ヘッドが集めた情報を混ぜ合わせ、元のベクトルに足せる形に整えるためのもの)を掛けたものが、Attention 層の出力(10 × d)です。

どのヘッドが何を担当するかは人間が決めるのではなく、学習の結果として分かれていきます(上の例のように、きれいに役割を言い表せるヘッドばかりではありません)。