Transformer 順伝播

21 個の primer が self-attention、multi-head attention、位置エンコーディング、 block、3 つのアーキテクチャ形態を別々に作ってきた。このページはそれを 1 つの機械として動かし、サンプリング、cache、実際の費用まで進む。

01

スタック全体の形

21 個の primer が部品を作った。ここではそれが 1 台の機械として動く —— 6 段、1 つの文、最初から最後まで。

段はいつも同じだ:トークン化、embedding と位置の付加、同じ block を N 個実行、正規化、逆埋め込み、logits の読み出し。attention は block の中に住み、形は前後の帳簿づけの中にある —— だからそこから始める。

トークンはまだ意味ではない —— それは id であり、その idは、各単語が 1 行を持つ表の行番号にすぎない。文をドラッグして、この検索の着地点を見てみよう:

位置 0 ·「the」· 表の 0 行目

位置 0 とはどちらも「the」で、まったく同じ行に着地することに注目してほしい —— 同じ id、同じベクトル、自分が文中の何番目かという情報は一切ない。これは意図的だ:順序はまだこの絵の中にない。それは次に登場し、位置エンコーディングというテーマそのものになる。

引いて見ると、6 段は 1 枚の図になる:アクティブな段が点灯し、その手前の段はすでに実行済みだ。順伝播全体を 1 段ずつ進めてみよう:

1 / 6 段 —— トークン化

block の帯がわざと高く描かれていることに注目してほしい —— それは 1 段ではなく、同一の block を N 個積んだものであり、N は次に回すつまみだ。self-attention、multi-head attention、位置エンコーディングが作ったものはすべて、あの 1 本の幅広い帯の中に収まっている。

そのつまみを回してパラメータ数を見てみよう —— 幅は GPT-2 small 自身の 768 に固定し、深さだけを動かす:

N = 12 · 124M パラメータ

GPT-2 small 自身の 12 層は 1.24 億パラメータに着地する —— 曲線からそのまま読める数字で、論文が報告している数字でもある。深さだけでそこに届いた:語彙も幅も変えず、 700 万パラメータの block を 12 個、共有の embedding 表の上に積んだだけだ。まで押すと、パラメータ数はおよそ 3 倍になる —— 幅は変わらず、曲線は直線のままだ。

その N 個の block はどれも残差ストリームを受け取り、できることは 2 つしかない —— 加えるか、置き換えるか。切り替えてみよう:

block 1 · 加える · stream の幅 4

「置き換える」が何をするか見てほしい:stream自身の幅は変わらない —— 依然として 4 スロット幅だ —— が、それ以前の block が書いたものは消えている。これが書き留める価値のある不変条件だ:どの層の境界でも stream は形を保ち、どのサブ層も加えるだけで置き換えない。§02 は block を開き、実際に何が加えられるかを見る。

02

1 つの block を、2 回見る

その N 個のコピーはどれも同じ 4 手を実行する。1 回たどれば終わりだ —— 仕組みは別の primer が持っている。

block は stream を読み、コピーを正規化し、その上で attention を実行し、結果を足し戻す。FFN でも同じことをする。どちらのサブ層も §01 の「加える、置き換えない」という形に包まれている。

アクティブな手は進めるたびに点灯する。真ん中の直線はstream そのもので曲がらない —— どの枝もそこから分かれ、そこへ戻る:

4 段中 1 —— 正規化

両方の枝が同じ形で終わることに注目してほしい:attentionが加え、次に FFN が加える。これは §01 の不変条件が 1 つの block の中で 2 回起きているということだ —— ここでは何も stream を置き換えない。

この 2 つのサブ層は同じ仕事をしていない。トークンを 1 つ選び、それぞれが何を読めるか見てみよう:attentionは他のどの位置も読める。FFNは自分が座っている位置しか読めない:

トークン 0 · attention は 4 個すべてを読む · FFN は自分だけ

トークン 0 から離れた瞬間、attention の線は 4 つのセルすべてに広がるが、FFN の線は自分の列から一度も出ない。はっきり言えばこうだ:attention はトークン間で情報を運び、FFN はチャンネル間で情報を運ぶ、一度に 1 トークンずつ。

その 2 手目はわざとコストが高い。1 つの block が持つ 12d² 個のパラメータ、幅をドラッグして、どう分かれるか見てみよう:

attention 2.4M(33%)· FFN 4.7M(67%)

FFN はどの block でも重みの 3 分の 2 を持っていく —— 各 block の隠れ層が d_model の 4 倍幅だから、その代金を 2 回払う、拡大で 1 回、縮小で 1 回。 を試してほしい:具体的な個数は縮むが、この比率自体は動かない。それは比であって、量ではない。

block にはもう 1 つ選ぶことがある:正規化をどこに置くかだ。切り替えて、勾配自身が embedding に戻る道を見てみよう:

N = 12 · pre-norm · 勾配は正規化を 0 回くぐる

post-norm の下 —— 2017 年の最初の設計 —— では、残差路自体がどの block でも正規化を 1 回くぐる。N 個つなげば N 回だ。pre-norm の下では正規化は枝にしか触れず、stream には決して触れないので、勾配がくぐる回数は 0 になる。この 1 つの選択が、GPT-2 も Llama も百層を超えて安定して学習できる理由の大半だ。§03 はこのスタックを離れ、反対側から出てくるものを問う。

03

logits からトークンへ

スタックの終わりは logits、語彙の各項目に 1 つのスコアだ。それを 1 語に変えるのは、それ自体が 1 つの小さな機械で、このページで初めてスタックが作っていないものだ。

物語のある一点を取ろう:prompt は「the cat sat on the ___」で、モデルは 5 個の候補 —— mat、rug、floor、sofa、roof —— に点数をつけた。それぞれの生スコアを 1 個ずつたどってみよう、何かがそれを作り変える前に:

「mat」· 生スコア 4.0

生の数字だけで見ても「mat」がすでに先頭で、4.0 が他を引き離していることに注目してほしい。softmax はどんなスコアの列も合計 1 の確率に変える。その分布の形を —— 先頭が入れ替わるわけではなく —— 実際に変えるのは温度だ。そのとき勝っている候補はスライダーを動かすたびに描き直される。静止状態の T = 1 では、すでに「mat」が 68.7% を占めている:

T = 1.00 ·「mat」68.7%. 左右にドラッグ。矢印キーで 1 段ずつ、Home で開始状態に戻る
T = 1.00 ·「mat」68.7%

T が 1 より上に上がるにつれ棒が平らになっていくのを、 T が 0 に向かって下がるにつれ「mat」がほぼすべてを飲み込んでいくのを見てほしい。ではすでにほぼ全部を持っていく —— これが極限での argmax の姿だ:常に最も高い 1 つの logit を、毎回。

argmax は裾を切り落とす 1 つの方法であり、上位 k 個だけを残すのは別の方法だ。 k を下げてドラッグし、切り捨てられる確率が増えるのを見てほしい:

k = 5 · 確率の 0.0% を切り捨て

ではこれは argmax そのものだ —— 候補 1 個、代案ゼロ。 k は分布の形にまったく適応しないことに注目してほしい:モデルがどれだけ自信を持っていようと、常に固定の個数へ切り詰める。

top-p は適応版だ —— 累積確率がちょうど p に届く最小の集合だけを残す。分布が尖っていれば残る候補は少なく、平らなら多く残る:

p = 1.00 · 5 個残す · 0.0% を切り捨て

ではちょうど 2 つの候補が生き残る —— mat と rugを合わせると 90% をわずかに超え、残りは入らない。p を下げれば集合は 1 個まで縮み、 p を 1 に近づければ、確率がわずかでもある候補はすべて残る。

貪欲デコード —— 毎ステップ argmax、それを永遠に —— には実在の失敗モードがある:自分の手で追い込んで見てみよう:

ステップ 0 ·「the」を出力

このおもちゃモデルの argmax 経路が自分に戻ると、貪欲デコードはそのトークンを永遠に繰り返す —— 気づく仕組みも抜け出す方法もない。正しい瞬間に 1 回サンプリングすれば循環は破れる。本番ではこれを repetition collapse と呼ぶ。§04 は、この節が選んだトークンを戻した瞬間に何が起きるかを問う。

04

1 トークンずつ生成する

次の 1 語だけを採点するモデルは、まだ生成器ではない。勝者を、最初からprompt の一部だったように戻せば生成器になる。

スタックを 1 回通るたびに系列は 1 トークン伸びる:順伝播し、サンプリングし、追加し、また繰り返す。「生成モード」という別物はない —— §01–03 で組み立てた同じ順伝播を、毎回長い入力へ呼び直すだけだ。

promptは与えられたもので、最初から灰色だ。生成トークンはその後ろに定着していく。スライダー下の 1 個は、モデルが今決めているものだ。1 段ずつ進めてみよう:

ステップ 4 · prompt ·「the」. 左右にドラッグ。矢印キーで 1 段ずつ、Home で開始状態に戻る
ステップ 4 · prompt ·「the」

定着したトークンはすべて次のトークンの入力の一部になることに注目してほしい —— これが自己回帰の定義そのものだ:ステップ t の出力は、ステップ t + 1 の入力の材料の 1 つになる。

注意しないと、これには代償がある。naive な方法は既存のトークンすべてを、すべての層で、毎ステップやり直す。賢い方法は新しいものだけを処理する。同じステップで比べてみよう:

現在 n = 5 トークン · naive · 6 個をやり直す

「naive」では、1 段進めるたびに前半部分がまるごと再点灯する —— n トークン分の仕事をして、トークン n + 1 を 1 個作る。「cached」では、最新のセルだけが点く。

その差を実際の生成にわたって走らせてみよう —— おもちゃ規模、4 層、d = 64 —— 2 つの合計がどう離れていくか見てほしい:

T = 12 · naive 51.05M · cached 4.85M

最初のトークンでは両者のコストはまだ近い —— 5 倍差だ。までに、naive の合計はcached の合計の 10 倍を超え、差はまだ開き続けている:naive のコストは生成量の 2 乗で増え、cached のコストは線形にしか増えない。

これを GPT-2 small 自身の 12 層・768 次元まで拡大し、20 トークンの promptから始めてみよう:

1 個生成 · naive 3.4e+09 FLOPs · cached 1.7e+08 FLOPs · 20.0倍

までドラッグすると、比は 69 倍を超えて着地する。これは丸め誤差ではない —— 本番システムが前半を実際に再実行することが決してない、その理由のすべてだ。代わりに何を、どれだけ保存しているのかが §05 だ。

05

cache が覚えているもの

「cached」が実際に保存しているのは、attention がすでに計算し終えたすべての key ベクトルと value ベクトルだ —— トークンごと、head ごと、層ごとに 1 組ずつ、二度と触れられない。

attention は新しいトークンの query を、それ以前のすべての key に対して採点する。その古い key と value は一度書かれたら変わらない —— それを生んだ block が二度とその上で走ることはない —— だから cache はそれらを正直に置いておくだけの場所だ。

生成されたトークンはそれぞれちょうど 1 組を追加する。埋まったスロットが積み上がるのを見てほしい:

1 組キャッシュ済み · 最新は「the」

これがこの節全体を支えている操作上の不変条件だ:トークン t を出力し終えた後、 cache はちょうど t 組を保持しており、トークン t + 1 を作るともう 1 組が追加される —— すでにそこにあるものは、再計算も上書きもされない。

この帳簿を GPT-3 自身が公表した形 —— 96 層、d_model = 12,288 —— まで拡大すると、もう無料ではなくなる。文脈をドラッグして伸ばしてみよう:

文脈 2,048 トークン · 9 GiB

GPT-3 自身の 2,048 トークン文脈では、たった 1 系列分の cacheだけで fp16 で 9 GiB に達する —— 重みを 1 つも読み込む前にだ。まで戻すと、ちょうどその 4 分の 1 になる:式は線形で、数値 1 個につき 2 バイト、key ベクトル 1 本と value ベクトル 1 本、層ごとに 1 組、それにウィンドウ内のトークン数を掛けたものだ。

cache は通常、無限のリストではなく固定サイズのバッファだ。生成をその容量を超えて押し進め、naive なリングバッファが何をするか見てみよう:

位置 0 → スロット 0 · 空き。左右にドラッグ。矢印キーで 1 段ずつ、Home で開始状態に戻る
位置 0 → スロット 0 · 空き

スロット 7 を過ぎると書き込みは折り返し、まだ有効な、もっと前のトークンが必要としているスロットに着地する。バッファの端を誰も守っていなければ、これは静かに失敗する:エラーもクラッシュもなく、古いトークンの key が新しいものに黙って置き換わるだけで、それ以降そこを参照して計算されるスコアはすべて間違ったものになる。

キャッシュされた各 key にはもう 1 つ、一緒に運ばれなければならないものがある —— 自分がどの位置にいるかだ。実在した実装バグを切り替えて、それがどうずれていくか見てみよう:

間違い —— 常に 0 · 書き込んだ位置 = 0 · 実際の位置 = 0

ここを間違える —— トークンごとに cache の現在の長さではなく常に position 0 を書いてしまう —— と、§01 からこのページが前提にしてきた位置エンコーディングと合わなくなる:モデルは新しいトークンを毎回最初の 1 個として採点し、出力は例外なく静かに劣化する。 §06 は、これを実際に動かす費用に数字を当てる。

06

予算

2 つの問い、どちらも実数の価値がある:1 回のパスに何がかかるのか、メモリは実際どこへ行くのか。

どのステップも重み行列全体を 1 回読み込み、そのパスのトークン数だけ使い回す。その比 —— 重み 1 バイト読むごとに得られる FLOPs —— が、そのステップが演算律速かメモリ律速かを決める。

prefill は prompt の全トークンを 1 パスで採点するので、1 回の重み読み込みが全員に行き渡るが、decode は 1 個にしか行き渡らない。prompt が伸びるにつれて 2 本の棒が離れていく様子を見てほしい:

P = 2 · prefill 2.0 FLOPs/バイト · decode 1.0 FLOPs/バイト

prefill の強度は prompt とともに上がっていく ——まで押すと、重み 1 バイトあたり 16 FLOP を稼ぐ。素直に演算律速だ。 decode の強度はモデルの規模に関わらず、ちょうど 1 に釘付けになる —— これはメモリ帯域律速で、演算を追加しても解決しない。

では別の問いだ。GPT-3 自身の規模で、固定された重みと1 系列自身の cacheを並べて量ってみよう:

文脈 2,048 · 重み 325 GiB · cache 9 GiB

フルの 2,048 トークン文脈であっても、1 系列の cacheは 325 GiB の重みの隣ではほんの薄い切れ端だ。リクエスト 1 個は安い。予算が変わるのはリクエストが 2 個目からだ。

重みは 1 回だけ支払われ、稼働中のすべてのリクエストで共有される。cache はそうではない —— 系列ごとにかかる。並列系列数をドラッグしてみよう:

1 系列 · 重み 325 GiB · cache 9 GiB. 左右にドラッグ。矢印キーで 1 段ずつ、Home で開始状態に戻る
1 系列 · 重み 325 GiB · cache 9 GiB

を超えて同時に保持し続けると、それらの合計 cacheは、その全員にサービスする重みよりも多くのメモリを食う —— まさにこの圧力が、生の FLOPs ではなく KV cache のメモリ管理こそを、本番のサービングシステムが実際に組み立てられている中心にした。

最後の数字で、この線を §01 と §02 まで巻き戻そう:GPT-2 small 自身の1.24 億パラメータのうち:

およそ 3 分の 2 が 12 個の block に住んでいる —— §02 の言い方をすれば、大半はFFN だ —— 残りは、このページが検索から始めたあの embedding 表だ。 §07 ではこのループを動くコードにまとめ、これらのコストを 1 つの表にまとめ、次に何を読むべきかを示す。

07

リファレンス

ループ全体を 9 行に、コストを 1 つの表に、そして 6 段をもう一度、それらを生成器に変える矢印とともに。

# grow the sequence one token at a time
cache = None
while len(tokens) < max_len:
    x = tokens[-1:] if cache else tokens
    logits, cache = model(x, cache)
    probs = softmax(logits[-1] / temperature)
    next_id = sample(probs, top_k, top_p)
    tokens.append(next_id)
    if next_id == EOS_ID: break

このループのどの行も、このページのどこかの節がすでに数字を当てたものだ —— 1 行ずつたどって、どの節か見てみよう:

model(x, cache) —— §01–02

この 4 行を合わせると実際いくらかかるのか:

時間
cached:新トークン 1 個あたり O(n)・uncached:O(n²)
n はここまでに存在するトークン数。attention 自体はどちらでも O(n) のままだ —— cache が省くのは、古いトークンをすべての射影と FFN でやり直す無駄な O(n·d²) であって、それらに対するスコア照会そのものではない。
空間
cache は O(layers · d_model · n)、数値 1 個につき 2 バイト
語彙サイズには依存しない —— 文脈長とモデル幅だけで増え、 n では変わらない重みの固定 O(params) の上に積み重なる。

この計算量が下がらない理由

天井になるのはめったに生の FLOPs ではない。prefill は演算律速だ。 decode はモデルの規模に関わらず重み 1 バイトあたり約 1 FLOP に釘付けになり、これは構造上メモリ帯域律速だということだ —— そして数十系列の並列を超えると、KV cacheはその全員にサービスする重みより重くなる。

バリエーション

prefillP トークン、1 パスO(P·d·layers) の cache を構築

ループを再びオンにすれば:

logits がトークン化へ戻る —— これが生成器たらしめる