Transformer Block 入門
ブロックとは、部品がついに出会う場所だ:attention と前向きネットワークが 1 つの関数の 2 つのサブ層として、残差接続が深さに勾配を消させないために、norm がその残差接続を本当に機能させる位置にちょうど収まり、そして全体を N 個、独立に積んでモデルにする。このページの数値はすべて隣の図から読み取れるもので、ページの締めくくりにある同じ 6 行の擬似コードから導ける。
2 つの仕事、1 つの形
attention はトークン間の情報を動かす、FFN は動かさない —— ブロックはこの 2 つを幅を変えずに走らせる。
このページのすべては 1 つの繰り返し単位、ブロックの中で起きる。言語モデルは同じブロックを何十個も積んだもので、各々が同じ形のベクトル列で返す。
ブロックでは 2 つの演算が順に走る:attentionはどの位置も他の全位置を読め、FFNは隣を見ることを許さない。切り替え token 2 に何が流れ込むか見てほしい:
FFN に切り替えると、収束していた線が 1 本の縦線に潰れることに注目してほしい:位置 2 に流れ込むのは位置 2 自身だけ。これが 2 つのサブ層の違いのすべてを 1 枚の絵にしたもの —— 位置を跨いで情報が動くのは attention の段だけで、それが先に走るので、後の FFN の手元には常に attention が集めた文脈がある。
1 つのトークンのベクトルがこの流れを通るのを追ってほしい:入ってきて、attention が更新し、FFN がもう一度更新する —— 3 段階をドラッグして、中の値だけでなく d_model 自体を見てほしい:
数値ではなく幅を見てほしい:6 本の棒が入り、attention を通っても 6 本、FFN を通っても 6 本。ブロック内のどの演算も d_modelをそっくり保つように作られている。次の演算 —— そしてその次のブロック —— との約束は「受け取ったのと同じ形を返す」、それだけだからだ。
この約束はどんな幅でも成り立つ。6 つの実モデルを滑らせて、からまで動かし、3 つのチェックポイントがずっと揃ったままなのを見てほしい:
これが attention と FFN を 1 つのものの「サブ層」と呼ぶ理由だ:両者は同じ入出力空間を共有し、ブロックは R^d_model から R^d_model への 1 つの関数になる —— この性質のおかげで、継ぎ目で形を合わせ直さずに N 個を積める。
スケールを正直に保つ
放っておくと、活性は深さとともに桁単位で漂う。正規化がその解決策だ ——LayerNorm と RMSNorm はちょうど 1 段だけ違う。
ブロックには、流れる数値に生来の下限も上限もない。各サブ層はベクトルを掛け、足し、重み付け直すだけで、範囲へ押し戻す仕組みはない。層数が十分あれば、これが積み重なる。
各層をベクトルの大きさに 1 つの固定係数を掛けるものとしてモデル化し、何も補正しないまま 24 層積む:その係数を 1 の少し上か少し下にドラッグして、 24 層後にどこへ着地するか見てほしい:
1.0 の両側にどれだけ余地が少ないかを見てほしい:1 層あたり係数 1.4 なら 24 層目で開始時の 3,000 倍超に達し、係数 0.6 なら 100 万分の 5 未満まで落ちる —— そして実際のブロックは 1 層あたり数十の演算を行い、そのどれもがこの係数を誰にも気づかれず 1 からずらしうる。
LayerNorm がその解決策で、8 次元のベクトル1 本に対しちょうど 2 段階:自分自身の平均を引き、自分自身の標準偏差で割る —— 1 段ずつ進んで、実際に計算される数値を読んでほしい:
計算された 2 つの数値がこの 1 本のベクトルだけに属すことに注目してほしい: μ と σ は各トークン、各位置、毎回独立に計算し直される。系列を跨いで共有されるものは何もない —— これこそ、バッチを必要とする BatchNorm と違い、長さの異なる入力の集まりでも LayerNorm が安全に使える理由だ。
RMSNorm は 2 段目だけを残す。同じベクトルの上でこの 2 つを切り替え、それぞれが残す平均を読んでほしい:
RMSNorm は平均を引かないので、出力は入力がたまたま持っていた平均をそのまま残す —— 破線はゼロから外れた位置にある。LayerNormの再中心化は、見た目ほど効いていないことになる:どちらで訓練しても到達する loss はほぼ同等で、RMSNorm の方が少ない演算でそこに着く。
どれだけ少ないか、具体的に:LayerNorm は平均、分散、再スケールを計算する。RMSNorm は平均のステップをまるごと飛ばす。6 つの実際の幅をドラッグして、 2 つの演算数を直接比べてほしい:
GPT-3 の幅では、その差はベクトル 1 本あたり数万回の要素ごとの演算になる ——のすべて、すべてのトークン、すべての順伝播と逆伝播で。これが「RMSNorm は全体でおよそ 7〜10% 速い」とよく引用される言葉の裏にある算術だ:原理的に仕事が少ないのではなく、呼び出し箇所ごとに実測で少ないのだ。
恒等はタダで手に入る
入力を出力に足し戻すと、ブロックが取りうる最も簡単な振る舞いは「何もしない」になる。この 1 回の足し算だけで、深いスタックが学習可能になる。
ベクトル x を f(x) に写す任意のサブ層を取り、後段へ渡すものを f(x) + x に変える。ブロックはもう出力全体をゼロから学ぶ必要がない —— 受け取ったものの上に補正だけを学べばいい。
残差接続をオフにしてまたオンにし、ブロックが後段へ何を渡すか見てほしい:
残差がオンのとき、すべて 0 を出力する f はブロック全体を恒等にすることに注目してほしい —— まさに y = x。そのようなブロックを積んでも全体としては何もしない。これは、他の何かを学ぶ前に「入力を忠実にコピーする」ことをまず学ばねばならないスタックよりも、遥かに親切な出発点だ。
既定値が恒等というのは半分だ。y = f(x) + x に連鎖律を使うとdy/dx = df/dx + 1、この+1が残差経路の貢献だ —— 各層の勾配を縮小係数としてモデル化し、勾配が渡る深さをドラッグしてほしい:
2 本の曲線がわずか数層で、徐々にではなく分かれることに注目してほしい:を過ぎると、プレーンな鎖はすでに、残差付きの鎖がまだ丸ごと運んでいるものの大半を失っている。第 1 層に事実上ゼロの強さで届く勾配は第 1 層の重みを更新できない —— その層はモデルには存在しても、学習からは不在になる。
この +1 は本物の第 2 の経路であって、同じ経路の明るさつまみではない:ブロックのスタックを動かし、skip 接続を「あり」と「構造ごと取り除いた」の間で切り替えてほしい:
レールは描かれるか描かれないかのどちらかなので、「少し弱くなった」と読み違える余地がない —— 取り除けば、どれだけ深く積んでいても第 1 ブロックに届く勾配はちょうどプレーンな鎖の数値になる。
これは Transformer に限った話ではない。残差接続が現れる前と後で、プレーンなスタックが実際にどこまで深く学習できたかを比べてほしい:
2015 年より前は、プレーンな畳み込みスタックに層を足しても約 20 層を超えると学習は悪化した —— より深いネットは、より浅いものにすら及ばなかった。その翌年、ResNet-152 は特別な工夫なしに学習できた。Transformer もまさにこの性質を継承している:GPT-3 の 96 層も Llama 70B の 80 層も、各層が残差ブロックだからこそ学習できる。
norm をどこに置くか
現代のブロックはみな、サブ層の前で正規化する、後ではなく。違いは見た目の話に聞こえる。だが深さが学習できるかどうかを決める。
2017 年の原論文は残差の和の後で正規化した ——post-norm。2020 年以降に作られたモデルはほぼすべて、代わりにサブ層の前で正規化する ——pre-norm。同じ 2 つの材料を並べ替えただけ —— だがこの順序が、残差接続が実際に何を守っているかを変える。
同じトークンを 2 つの配置に並べて通してみる:
最後の段がどこに着地するか見てほしい:pre-norm の残差の和はブロックの生の出力そのもので、norm には一切触れられていない。post-norm の和はすぐその場で正規化される、1 段遅れて。skip 接続はどちらも同一 —— 違うのは、それとブロックの出口の間に何か挟まっているかどうかだ。
この 1 段の違いが効くのは、逆伝播がブロックを逆向きに走るからだ:各 norm の逆伝播を、通過する勾配のうち割合 cだけを残すものとしてモデル化し、深さを増やしながら pre-norm とpost-norm を比べてほしい:
post-norm の +1 は norm の内側にあるので、1 層ごとに c が掛かる —— c を下げると、post-norm の曲線は残差接続ごと、プレーンな鎖がずっと持っていたのと同じ消失の形に折れ戻る。pre-norm の +1 は何の内側にもない。スタックがどれだけ深くても、c がどれだけ小さくても、常にちょうど 1 のままだ。
pre-norm にはまだ払っていない代償が 1 つある:ブロックの内側には残差ストリーム自体を縮め直すものが何もない。自身の大きさがブロックごとに増えていくのを、補正が見当たらないまま見てほしい:
norm はサブ層へ入る途中で x のコピーを読むだけで、直通路上の x には触れない —— これが、pre-norm のモデルすべてが下流の何かがストリームを直接読む前に、最後のブロックの後にもう 1 回 norm を必要とする理由だ。
これが、ブロックが両方を運ぶ理由だ:正規化は流れのスケールを、残差接続は流れが第 1 層まで届くかを制御する。norm を和の逆側に置けば、1 つ目をこなしながら静かに 2 つ目を損なう —— モデルが深くなるまで誰も気づかない。
完全なブロックを、一気に
部品を順に並べる —— norm、attention、加算、norm、FFN、加算 —— それが現代の decoder が走らせる pre-norm ブロックだ。
2 サブ層、2 残差接続、2 norm、常にこの順序で 1 回ずつ —— 前の 3 節が積み上げてきた組み合わせを、1 つのブロックで 2 回使う。
1 トークンのベクトルx が読まれ更新されまた読まれるのを、6 段階で見てほしい:
x が置き換えられるのではなく再利用されることに注目してほしい:同じ変数が連続して正規化され、変換され、また加算される、2 回。残差接続の視点からは、attention と FFN はそれぞれ累計値に小さな補正を加えているだけで、どちらもブロックの出力をゼロから作ってはいない。
この 2 つの補正のコストはまるで違う。6 つの実際の幅をドラッグして、 1 ブロックあたり attention が使う分と FFN が使う分を比べてほしい:
FFN の棒がどの幅でも常に attention の棒のちょうど 2 倍であることに注目してほしい:4 つの d_model×d_model 行列対 2 つの d_model×4d_model 行列は 4d² 対 8d²で、この比は d を完全に打ち消す。FFNがたまたま大きいのではない —— 幅によらない正確な倍率で大きいのだ。
モデル全体まで引いて見ても、同じ模様が別の尺度で繰り返される。 embedding テーブルと全ブロック合計を比べてほしい:
直感的には「attention こそモデル」と感じるが、1 ブロックのパラメータの大半を占め、トークンレベル計算の大半を担うのは FFN だ —— attention は情報をどこへ送るかを決め、FFN は届いた情報で何をするかを決める。この配分ゆえ、研究者はモデルの事実的記憶はFFN に宿ると考えるようになった。
組み合わせは深さで起きる
このブロックを N 個積んでも、どれか 1 つを見れば何も変わらない。変わるのは、同じ変換が自分自身と何回合成されたかだ。
各ブロックは独自の重みを持つ —— attention、FFN、norm がそれぞれ独自だ —— だがR^d_modelを読み書きする。これが §01 の形の契約だ。スタックを 1 ブロックずつ、実モデルの深さまで積んでほしい:
N が増えても図の中の何も形を変えないことに注目してほしい —— 変わるのはシルエットの数だけ:、。深さも幅もモデルを大きくする —— 共有の倍率をスライドさせ、N だけを増やす場合とd_modelだけを増やす場合を比べてほしい:
直線と曲線が引き離れていくのを見てほしい:深さを 4 倍にするとパラメータはちょうど 4 倍、幅を 4 倍にするとちょうど 16 倍になる。ブロック内のどの行列も d × d か d × 4d だからだ ——幅は 2 乗のコスト、深さは線形のコストで手に入る。
つまり同じパラメータ予算でも、どこに使うかでまるで違うものが買える。総量をおおよそ固定したまま、深く狭いスタックと浅く広いスタックの間をドラッグしてほしい:
同じ予算で狭いスタックがどれだけ多くのブロックを買えるか見てほしい ——幅の 2 乗コストこそが、深さを「もっと逐次的な組み合わせを買う」安い方法に、幅を「1 つの深さでもっと容量を買う」高い方法にしている理由だ。
「もっと組み合わせ」は比喩だけではない。N 個の残差ブロックを展開すると、信号が通りえたブロックの部分集合すべてについての和になる —— N をドラッグして、暗黙の経路数がどれだけ速く増えるか見てほしい:
で、すでに 1,680 万通りの経路になる、ごく普通の重みのスタック 1 つの中に(Veit、Wilber & Belongie、 2016 年)。幅はこれを生まない:ニューロン数を倍にしても 1 つの深さの容量が倍になるだけだが、層数を倍にすると組み合わせ数は掛け算で増える。
6 行と、それがなるモデル
このページの主張は、2 行を N 回繰り返し 4 行で包んだものに帰着する。
この包みを 1 段階ずつ、ブロック自体から完全なモデルまで組み立ててほしい:
合わせて 6 行だ:§05 が段階ごとに追ったあの 2 行のブロックを、今組み立てたスタックの中で N 回繰り返し ——
x = x + attn(norm(x)) x = x + ffn(norm(x))
さらに 4 行で包む —— 片端に embedding、もう片端に norm と head を 1 回 —— それで完全な decoder-only 言語モデルになる:
x = embed(ids) for blk in blocks: x = blk(x) x = final_norm(x) logits = head(x)
head は単一の d_model × vocab_size 線形層だ —— embedding テーブルと行列を共有し、パラメータを二重に払わずに済むことが多い。
ブロックの中で 1 つだけ自己注意 primerに残した部分がある:因果マスクがある位置に実際どこまで読ませるか見てほしい:
は登場した瞬間に 6 つ全部を読めるが、は自分しか読めないことに注目してほしい。推論時には各位置の K と V がキャッシュされ、再計算されない —— これが生成が長くなるほど続きが遅くなる理由だ。