コンテンツにスキップ

4. Transformer ブロック全体

この回では、ここまでの部品がどう組み合わさって 1 つのモデルになっているかを見ます。新しく出てくる部品は、残差接続・MLP・LayerNorm・出力層の 4 つです。

この回の入力と出力
  • ブロック 1 つの入力と出力: ベクトルの列(10 × d)→ 同じ形のベクトルの列(10 × d)。形が変わらないので、何十個でも直列につなげられる
  • モデル全体の入力と出力: トークン ID の列(10 個)→ 次のトークンについての、語彙(ボキャブラリ)トークナイザが知っているトークンの一覧。LLM は、この一覧のどれか 1 つを次のトークンとして選ぶ。語彙の数を語彙サイズという。用語集でくわしく →のすべてのトークンのスコア(ロジット(logits)モデルが最後に出す、語彙のトークン 1 つ 1 つに対するスコア。大きいほど「次に来そう」という意味だが、まだ確率ではなく、負の値もとる。用語集でくわしく →。語彙数ぶん。GPT-2 では 50,257 個)

Transformer は、同じ形の ブロック を何十層も積み重ねたものです。1 つのブロックは、次の 2 段構えになっています。

入力 x(10 × d)出力 x(10 × d)残差ストリームLayerNormAttention更新分を足す+LayerNormMLP更新分を足す+
左の太い縦線が残差ストリーム(x の流れ)。右へ分かれた枝(LayerNorm → Attention、LayerNorm → MLP)の出力が、+のところで本線に足されます。どの矢印の上でも、データの形は 10 × d です。

図では、左の太い縦線が xx の流れです。そこから右へ枝分かれして LayerNorm → Attention(または MLP)を通り、+のところで元の流れに足されて戻ります。

  1. Attention 層: トークン同士で情報をやり取りする(前回の内容)

    入力 xx(10 × d)を LayerNorm(層正規化)1 つのトークンのベクトルの数値を、平均が 0、分散(数値の散らばり具合)が 1 になるように揃え、学習した倍率を掛けて、学習したずれを足す処理。値が大きくなりすぎたり小さくなりすぎたりするのを防ぐ。用語集でくわしく → で整えてから Attention に通し、その出力を 元の xx に足します。

    x←x+Attention(LayerNorm(x))x \leftarrow x + \mathrm{Attention}(\mathrm{LayerNorm}(x))

  2. MLP 層: 各トークンが、それぞれ自分のベクトルだけを見て変換する

    同じように LayerNorm → MLP(多層パーセプトロン)線形変換と活性化関数を交互に重ねた、基本的なニューラルネットワーク。Transformer では、各トークンのベクトルを 1 つずつ別々に変換する部分を指す。用語集でくわしく → に通し、結果を足します。

    x←x+MLP(LayerNorm(x))x \leftarrow x + \mathrm{MLP}(\mathrm{LayerNorm}(x))

足し算をしているので、どの段階でも xx の形は 10 × d のままです。

残差接続:「書き換える」のではなく「書き足す」

Section titled “残差接続:「書き換える」のではなく「書き足す」”

上の式で大事なのは、各層の出力を 元のベクトルに足している(残差接続層の出力で入力を置き換えるのではなく、入力に出力を足し合わせるつなぎ方。x ← x + 層(x) と書ける。入力をそのまま伝える通り道ができるので、元の情報が残りやすく、深く積んでも学習しやすくなる。用語集でくわしく →)点です。

ベクトル xx は、入力から出力まで 1 本の流れ(残差ストリーム)として通り抜けます。各層は、その流れを置き換えるのではなく、「足し込む更新分」を出します。更新分は新しい情報を書き足すこともあれば、既存の成分を強めたり打ち消したりすることもありますが、元のベクトルが通る道は常に残ります。

前の回の小さな数値例で見てみます(説明を簡単にするため、LayerNorm と、ヘッドの出力にかける最後の行列は省きます)。「それ」のベクトルを x=[1,1]x = [1, 1](前の回で [1, 1] だったのは k ですが、ここでは元のベクトル xx もたまたま同じ値だったとします)、Attention の出力を [0.28,1.72][0.28, 1.72] とすると、足したあとは x=[1.28,2.72]x = [1.28, 2.72] です。もし置き換えていたら [0.28,1.72][0.28, 1.72] になり、元の [1,1][1, 1] は消えます。この例では、足すことで元の値の上に「食べ物」の成分が上乗せされています。

足し算でつなぐと、次の 2 つの良いことがあります。

  • 層が深くなっても、元の情報が失われにくい
  • 学習大量のデータで予測を試し、外れた分だけパラメータを少しずつ調整して、予測を当たりやすくしていくこと。LLM は「次のトークンの予測」で学習する。用語集でくわしく →のときに、勾配各パラメータを少し動かしたとき、予測の外れ具合(損失)がどれだけ増えるか減るかを表す値。学習では、勾配の逆向きにパラメータを動かして損失を減らす。用語集でくわしく →が、出力側から入力に近い層まで、何十層さかのぼっても届きやすい。何十層も積めるのはこのおかげ(上の図の太い縦線は、出力から入力まで何も掛けずに通れる「素通りの道」です。学習の手がかりは出力側から層をさかのぼって伝わりますが、層を通るたびに小さな係数が掛かると、何十層で消えるほど薄まってしまいます。素通りの道があれば、薄まらずに戻ってこられます)

MLP は、各トークンのベクトルを 1 本ずつ変換する小さなニューラルネットワークです。同じブロックの中では、すべてのトークンに同じ MLP(同じ行列)を使いますが、トークンどうしのベクトルは混ぜません。GPT-2 型の MLP は、次の 3 ステップです。

  1. 線形変換ベクトルに行列を掛けて、別のベクトルに変えること。各出力は、入力の数値に重みを掛けて足し合わせたものになる。用語集でくわしく →で、d 次元を 4 倍ほど(4d 次元)に広げる
  2. 活性化関数(GELU など)ニューラルネットワークの層の間に挟む、曲がった形(非線形)の関数。これがないと、線形変換を何層重ねても 1 回の線形変換と同じになり、複雑な関係を表せない。用語集でくわしく →(GELU など)を通す。ここで非線形線形ではない関係。グラフにすると 1 本の直線(平面)にならない。たとえば ReLU(負なら 0、正ならそのまま)は、ReLU(1) + ReLU(-1) = 1 なのに ReLU(1 + (-1)) = 0 で、「足してから変換する」と「変換してから足す」の結果が食い違う。用語集でくわしく →な変換が入る
  3. もう一度線形変換して、d 次元に戻す

広げた各次元は、「この特徴があるか」を調べる検出器のように働きます。一度広げるのは、検出器をたくさん用意するためです。前の節の「それ」のベクトル [1.28, 2.72]([動物, 食べ物])を、4 つの検出器で調べてみます。

  1. 広げる:4 つの検出器[動物, 食べ物, 食べ物 − 動物, 動物 − 食べ物]を計算する → [1.28, 2.72, 1.44, −1.44]。2 個だった数値が 4 個に増えた
  2. ReLU:負の値を 0 にする → [1.28, 2.72, 1.44, 0]。「動物寄りか」を調べる 4 つ目の検出器は反応しなかった
  3. 戻す:2 次元に戻す。ここでは 3 つ目(食べ物寄りか)の値の 0.25 倍だけを「食べ物」の次元に入れるとする → 更新分 [0, 0.36]。残差接続で足すと [1.28, 3.08]

この MLP は、「食べ物寄りなら、食べ物らしさをさらに強める」という規則を、行列の値(0.25 など)として覚えていることになります。本物の MLP には、こうした「検出器」と「書き足す先」の組が何千もあり、下で見る「事実の記憶」も、同じ形で蓄えられていると考えられています(MLP をこうした組の集まりとして読む見方は、Geva ら、2021 年 などによる解釈の 1 つです)。

(この例は説明用です。LayerNorm は前の節と同じく省き、活性化関数は GELU の代わりに、負の値を 0 にするだけの ReLU を使い、本来は 4 倍の 8 次元に広げるところを 2 倍の 4 次元にしています。)

Attention が「他のトークンから情報を集める」のに対し、MLP は「集めた情報をもとに、自分の中で考える」部分だと言えます。 「マイケル・ジョーダン」というトークン列から「バスケットボール」という知識を引き出すような、事実の記憶の多くは MLP に蓄えられている と考えられています(3Blue1Brown の動画「LLM はどう事実を記憶するのか」 がこの話です。研究が続いている分野で、すべてが解明されているわけではありません)。パラメータモデルが学習で調整する数値。埋め込み行列や、Attention・MLP の行列の中身がすべてパラメータ。「70B のモデル」は、パラメータが 700 億個あるという意味。用語集でくわしく →の数で見ても、各ブロックの約 3 分の 2 は MLP が占めています(4 倍に広げる場合)。

LayerNorm:値のスケールを揃える

Section titled “LayerNorm:値のスケールを揃える”

LayerNorm は、各トークンのベクトルの数値を、いったん平均が 0、分散(数値の散らばり具合)が 1 になるように正規化数値の大きさや範囲を、扱いやすい基準に揃えること。たとえば「平均 0・分散 1 にする」「合計を 1 にする」「長さを 1 にする」など。用語集でくわしく →し、そのあと学習した倍率とずれで調整する処理です。Attention や MLP に入れる値の大きさを揃え、層を重ねても各層の計算が極端な値で不安定にならないようにします。

上の図のとおり、LayerNorm がかかるのは枝に入る値だけで、左の本線の xx にはかかりません。本線には層ごとに更新分が足されていくので、最後に出力層へ渡す前に、もう一度 LayerNorm で揃えます(下の「モデル全体」の 4.)。

[2, 4, 6] というベクトルなら、次の順に計算します。

  1. 平均(4)を引く → [−2, 0, 2]
  2. 標準偏差(散らばりの大きさ。分散の平方根)で割る。分散は 1. の値の 2 乗の平均で (4 + 0 + 4) ÷ 3 ≈ 2.667、標準偏差は √2.667 ≈ 1.633 → 約 [−1.22, 0, 1.22]
  3. 次元ごとに学習した倍率を掛け、次元ごとに学習した値を足す(揃えすぎた値を、モデルが使いやすい大きさと位置に戻せるようにするため)

Attention と MLP を 1 回ずつ通すだけでは足りず、ブロックを何十層も重ねるのには、主に 3 つの理由があります。

1. 情報を「伝言リレー」で運ぶには、段数が要る

Section titled “1. 情報を「伝言リレー」で運ぶには、段数が要る”

1 つの層の中では、「A の情報を B が受け取り、その B を C が見る」という伝言はできません。 各層の Attention が見るのは、どのトークンについても、その層に入ってくる時点のベクトル(前の層までの結果)だからです。その層で誰かが集めた情報は、次の層になって初めて、ほかのトークンから見えるようになります。だから、伝言を回すには層を重ねる必要があります。

例文で、「それ」と「だった」のベクトルに何が足されていくかを、層ごとに追ってみます(中身はたとえで、実際のモデルがちょうどこの層でこう動くとは限りません)。

「それ」のベクトル「だった」のベクトル
埋め込みの直後「それ」という語だけ「だった」という語だけ
1 層目の後前の「食べた」などの文脈が薄く混ざる。まだ、どれを指すかは決まっていないが、「食べられたものを探そう」という手がかりになる「新鮮」「それ」などを見て、「何かが新鮮だった」の情報が入る
2 層目の後その手がかりで q が食べ物の向きになり、「魚」に強く注目して「それ=魚」の情報が入るまだ「それ=魚」は届いていない(1 層目の後の「それ」は、まだ指す先を特定していなかった)
3 層目の後(さらに情報が足される)2 層目で「魚」を知った「それ」を見るので、「新鮮なのは魚」という情報が届く

「それ=魚」という情報が「だった」に届くまでに、3 層かかりました。前の段で得た結果を使って次を調べる、という処理を多く重ねるほど、たくさんの段が要ります(文章が長いだけで、必要な段数が増えるわけではありません)。

2. 「集める → 考える → それを踏まえて、また集める」ができる

Section titled “2. 「集める → 考える → それを踏まえて、また集める」ができる”

Attention は情報を 集める 係、MLP は集めたものを 加工する 係です。MLP が加工した結果は残差ストリームに足され、次の層の Attention で Query(何を探すか)や Key(自分が何を持っているか)を作る材料になります。

人が文章を読むときも、「『それ』は何を指す? → 魚だ → では『新鮮だった』は何の話?」と、分かったことを足場にして次を調べます。1 回集めて 1 回考えるだけでは、この足場を使った読み方ができません。

3. 深くすると、少ない部品で複雑な処理を表せる

Section titled “3. 深くすると、少ない部品で複雑な処理を表せる”

段を重ねると、単純な特徴を組み合わせて、複雑な特徴を作れます。同じ処理を 1 段だけで表そうとすると、ずっと多くの部品(上の MLP の例の「検出器」のような、数値を受け取って 1 個の数値を返す小さな計算。ニューロンと呼ぶ)が要ることがあります。

ただし、深くするほど計算は重くなります。何層にするかは、精度と計算のコストの兼ね合いで決まります(GPT-2 最小構成で 12 層、大きなモデルでは数十〜百数十層。下の「モデル全体」)。各ブロックの形は同じでも、パラメータは別々なので、同じ処理を繰り返しているわけではなく、層ごとに違う役割を学習します。

  1. トークン化して ID の列にする(10 個)

  2. 埋め込み行列でベクトルに変換する(10 × d。GPT-2 などでは位置の情報も足す)

  3. Transformer ブロックを NN 回通す(10 × d のまま。GPT-2 最小構成で 12 層。大きなモデルでは数十〜百数十層で、たとえば Llama 3 の 70B で 80 層、405B で 126 層、GPT-3 の 175B で 96 層)

  4. 最後の LayerNorm を通す(10 × d)

  5. 一番最後のトークン(例文なら「だった」)のベクトル 1 本(d 次元)に、「d × 語彙数」の行列を掛けて、語彙のすべてのトークンのスコア(logits、語彙数ぶん)を出す。この最後の変換を 出力層 と呼びます。GPT-2 では、この行列に埋め込み行列の転置行列の行と列を入れ替えること。A の転置を Aᵀ と書く。用語集でくわしく →を使い回します

    たとえば説明用に d = 2、語彙が「。」「ね」「が」の 3 個だけだとします。最後のベクトルが [2, 1]、行列の各候補の列が「。」[1, 0]、「ね」[0, 1]、「が」[1, −2] なら、スコアは各列との内積で、2×1 + 1×0 = 2、2×0 + 1×1 = 1、2×1 + 1×(−2) = 0 → logits [2, 1, 0]。次の回の softmax の例でも、この [2, 1, 0] を使います。

最後のトークンは、Attention で前のすべてのトークンを見られます。学習によって、そこから予測に役立つ文脈を取り込むようになっているので、「次に何が来るか」の予測に使えます。 (実際には全トークンについて同じ計算ができ、学習のときは「各位置で次のトークンを当てる」問題を一度にまとめて解いています。生成のときに使うのは最後の 1 つだけです。)

最後の logits から次のトークンをどう選ぶかは、次の回で扱います。