LLM のテキスト 入門
モデルとは浮動小数点数の長方形の上の関数だ。文字列は長方形ではない。この primer は 2 つのあいだの道のりと、テキストにあって長方形にはない 3 つの性質を扱う:長さが可変であること、順序そのものが意味を運ぶこと、そして各トークンの意味が隣で決まること。ここの数値はすべて隣の図が計算している。
モデルはテキストを見ていない
見ているのは数値の長方形だ。面白いことはすべて、そこへ至る途中で起きる。
Transformer のあらゆる演算 —— 行列積、softmax、残差加算 —— は固定された長方形に並んだ浮動小数点数の算術だ。文字列はそのどれでもない。両者のあいだには 5 回の変換があり、どれが可逆でどれが黙って何かを捨てているかは知っておく価値がある。
そこで短い文字列 café ☕ を最後まで運ぶ。スクラバーを 1 つ進めるごとに 1 段階、いま立っている段階の個数がその横に出る:
個数がほぼ毎段階で変わり、しかも 2 度と同じ意味にならないことに注目してほしい。6 文字が 6 コードポイントになり、9 バイトになり、3 トークンになり、3 個の ID になり、最後に 3 行の浮動小数点数になる。モデルが実際に食べるのは最後のものだけで、そのとき元の文字列から残っているのはID だけだ。
下から見ていこう。UTF-8 は 1 コードポイントに 1〜4 バイトを使う。この図はそれを文字どおりに描く:文字の箱の幅は、その下のバイトの数とぴったり同じだ。スライダーで 6 つの文字体系をたどってほしい:
箱が等幅でなくなるのを見てほしい。英語は 1 文字 1.00 バイト、、、絵文字は 4.00。 4 対 1 の差があり、どれを払っているかを告げるものは何もない。ASCII が 1 バイトのままなのは意図的で、だから古いテキストはそのまま正しい UTF-8 になる。
いまの段落の「文字」という語は働きすぎだ。読者が1 文字と呼ぶものは複数のコードポイントが貼り合わさったものでありうるし、プログラミング言語が返す数は 3 つのどれでもない。6 つの例を滑らせてみよう:
len() が数えるのはコードポイントなので、 ——1 グリフ、18 バイト —— は5 と答える。任意の位置で文字列を切れば、文字の途中を断ち切ることになる:その絵文字に s[:1] をすると男が 1 人だけ残り、バイトで切れば正しい文字にすらならない。
さらに悪いことに、1 つのグリフに正しい綴りが 2 通りありうる。NFC はアクセントを文字に合成し、NFD は文字を残して結合記号を後ろに吊るす。形式を切り替えると、見た目は動かないままバイトだけが動く:
どちらの形式も正準等価で、どちらも正しく、例外も出ない。合成形では 5 バイト、分解形では 6 バイトで、3 バイト目から食い違う —— つまり "café" == "café" は False であり、下流のハッシュ・索引・完全一致フィルタは、読者に1 つの語に見えるものを 2 つ抱えることになる。
この失敗は決して自己申告しない。キーがページ上に堂々と見えているのに、検索が何も返さない、という形で現れる。書き込み経路に正規化を押し上げて、誰にも見えないバケツが空になるのを見てほしい:
すべてのキーが書き込み時に正規化されればバケツは 2 つでなく1 つになり、クエリは 8 つすべてを見つける。それまでは、読者には区別できない重複がストアに残り、唯一の症状は再現率が本来より少し低いことだけだ。正規化は境界で行う:トークン化のあとでは遅い。そのときには 2 つの形式はもう別の ID になっている。
長さは形ではない
「OK。」は 2 トークン、Wikipedia の記事は 100 万トークン。同じ行列積が両方を受け取らなければならない。
GPU が欲しいのは長方形だ:B 行、T 列、どの行も同じ長さ。テキストには自然な上限も下限もないので、誰かがそれを長方形にしなければならない —— そしてそのやり方ごとに請求書が違う。
普通の答えはパディングだ:バッチの中で最も長い文書を取り、残りの行をパディングセルでそこまで埋める。12 番目の文書を長くドラッグしてほしい:
たった 1 行のせいで長方形全体が育つのを見てほしい。最短のときでもバッチはセルの 32% を無駄にし、外れ値を まで引くとテンソルの 69% は何も持たない。そのセルのひとつひとつが、本物のトークンとまったく同じように掛けられ、softmax され、加えられる。
無駄な算術は、モデルが結果を無視すると知っているあいだは耐えられる。それがマスクの仕事のすべてで、古典的なバグはそれに尋ねないプーリングだ。文書を短くし、平均が何で割るかを切り替えてほしい:
マスクなしの平均はノイジーなのではなく誤りで、しかも一方向に誤る:幅 12 の行に本物のトークンが 5 個なら0.300 と読むが、真の値は0.720 だ。例外は出ない。モデルは学習し、損失は下がり、ある指標だけが理由の分からないまま数ポイント低い。
無駄も自然法則ではない —— 誰が誰と長方形を共有するかの結果でしかない。コーパスを長さで並べ替え、バケツに切り、各バケツを自分の最長にだけ揃えてみよう:
1 バケツが素朴なバッチだ:本物のトークン 99 個に 220 セル、パディング 55%。にすると 110 セル、10% になる。落とし穴は反対の端にある —— 1 文書 1 バケツは無駄ゼロだがバッチサイズ 1 で、それはこの長方形を買った目的そのものだ。
もう 1 つの答えは尾を切り捨てることだ:最大長を決め、それより先を全部落とす。4,000 トークンの文書に切れ目を通してドラッグしてほしい:
では、モデルは4,096 のうち 512 を読み、残りの 88% があったことを決して知らない。エラーにはならず、答えを返す。この節のすべての失敗はこの形をしている:テンソルは正しく、算術は走り、間違っているのは数値の意味だけだ。
2 まで数えられない袋
3 語の 6 通りの並び、ベクトルは 1 つ。ベクトルしか見ないものには区別できない。
テキストを固定の形にする最も安い方法は、語がどこにあるかを気にするのをやめることだ:語彙の各項目が何回現れたかを数え、その計数を渡す。1 パス、長さは自由、しかもスパム判定やトピック分類では今でもまともなベースラインだ。
同時にそれは、言語を言語たらしめているものを捨てている。3 語の文のすべての並びをスライダーでたどり、下の計数ベクトルが動かないのを見てほしい:
6 通りの並び、ベクトルは 1 つ —— そしては「犬が猫を噛む」ではない。これはデータで良くなる近似ではない:系列から袋への写像は単射ではないので、袋だけを見る関数はどんな規模でも 2 つを分けられない。
古典的な手当ては、単語ではなく連続する短い並びを数えることだ。窓を広げて、モデルが何を数えられるようになるか見てほしい:
でモデルはようやく「犬噛む」と「噛む犬」を区別できる。下の組がその証拠だ。だが右側の数を見てほしい:特徴空間は語彙の n 乗なので、窓が 1 語広がるたびに可能性が 50,257 倍になる。
この取引は軸の上で見る価値がある。縦軸は対数で —— グリッド線 1 本ごとに 10 万倍 —— なので指数的な爆発は直線として描かれる:
軸が対数だからこそ、その直線が爆発そのものだ: では1.27×10¹⁴ 通りの 3 つ組がありうる。どんなコーパスでも埋めきれない。ほぼすべての特徴はゼロで、残りの多くも 1 回しか現れない。
しかもまだ遠くまで届かない。ならばをもしから引き離し、両方を含む窓が 1 つずつ消えるのを見てほしい:
距離が窓の幅に追いついた瞬間、両方を含む窓は存在しなくなる —— つまりモデル全体のどの特徴もその対に言及しない。もっと数えても助けにはならない。その対はそもそも特徴空間にないからだ。これが壁であり、この primer の残りが計数ではなく距離の話になる理由だ。
トークン単体に意味はない
川辺の bank と金利を決める bank は、同じ ID で同じ表の行だ。
算術が始まる前に、各 IDはベクトルにならなければならない。最も安い方法は表引きだ:ID 31 は行 31 を意味する。いつでも、どこでも、周りの文が何を言っていようと。
下の文に 2 回出てくる bank は同じトークンなので、どちらも同じ行を読む。スライダーで 2 つを行き来し、矢印の落ち先を見てほしい:
動くのは矢印だけだ。Word2Vec と GloVe はまさにこの表で、one-hot からの飛躍は巨大だった —— だがベクトルは文を読む前に決まっているので、 2 つの bank を分けるものは下流の別の何かが取り戻すしかない。
取り戻すとは、隣を混ぜることだ。α を上げて、 1 行が 2 行になるのを見てほしい:
α = 0 では 2 行は同一で、読み出しは距離 0.00 と言う。 では1.12 離れ、モデルは「どちらの bank か」に答えられる。この混合こそアテンションが計算するもので、重みはトークンごと・層ごとに学習される。
再帰型ネットは 1 歩ずつ混ぜるので、どれだけ前まで運べるかに固い天井ができる。ゲートを決め、距離をドラッグしてほしい:
1 歩あたり 90% を残すゲートでは、のトークンは状態の1% 未満だ。γ を上げれば届く距離は伸びるが、状態が飽和する危険も増す。0.80 まで下げれば 20 トークンで 1% をわずかに超える程度になる。段落を覚えつつ安定する設定は存在しない —— これが勾配消失を順方向から見たものだ。
2 つのトークンの隔たり
ここまでの構造はすべて 1 つの数で採点できる:信号がある位置から別の位置へ渡る最短経路の長さだ。
経路が短ければ長距離の構造を学べる。長ければ学べない。途中で通るものすべてが、それを上書きしてよいからだ。
アーキテクチャを選び、2 つを引き離してほしい。階段の 1 段が、信号の 1 跳に当たる:
袋には経路がそもそもない。3-gram は距離 2 までは 1 跳、それより先はない。カーネル 3 の畳み込みには ⌈d/2⌉ 層が要る。には 12 回の逐次ステップが要る —— 1 段が乗算 1 回、あの減衰の図そのものだ。アテンションはどの距離でも 1 跳で済む。
どこでも 1 跳、はただではない。すべてのトークンからすべてのトークンへ経路を持たせるには、すべての組に点を付ける必要がある。n をドラッグし、行ではなく正方形を見てほしい:
系列を倍にすると行は倍、正方形は 4 倍になる: なら 12 トークンに対して 144 のスコアで、そのすべてが計算され、softmax を通り、掛け合わされる。経路長は一定で、その代金を払っているのは面積だ。
この長方形の値段
系列がモデルの幅より長くなるまで、アテンションは安いほうの層だ。その交点は謎ではない。
幅 d の自己アテンション層は 2n²d 回の積和を行い、同じ幅の再帰層は 2nd² 回を行う。割り算すれば比は n/d なので、ちょうど n = d で同じ値段になる。
下の図はどちらの軸も対数なので、べき乗則は直線として描かれ、交点は 2 本の線が出会う場所そのものだ。系列長をそこまで引いてほしい:
GPT-2 の幅である で、アテンションと再帰型の曲線はどちらも 9.06×10⁸ 演算で交わる。左側ではアテンションが安く、かつ浅い。右側では、2 倍の傾きの線と引き換えに 1 跳の経路を買っている。
計算量はよく引かれるほうの半分だ。もう半分は、softmax が走るあいだスコアがどこかに存在しなければならない、ということだ。文脈長をドラッグしてメモリを読んでほしい:
で、12 層 12 ヘッドのモデルは 6.04×10⁸ 個のスコアを抱える ——fp16 で 1.13 GiB。 1 系列ぶんで、重みはまだ 1 つも数えていない。この数値が FlashAttention の存在理由だ:同じ softmax をタイルに分けて計算し、行列を書き出さない。あの 2 乗は、計算の問題になるはるか手前で、まずメモリの問題なのだ。
黙って失敗する 4 つ
形の誤りは例外を投げる。この 4 つは投げない。
正規化。2 つのバイト列、1 つのグリフ。索引は両方を保持し、再現率の低下は誰にも帰属できない。マスク忘れ。パディング後の幅で取った平均はもっともらしく、パディングの割合だけゼロ側へ偏る。切り捨て。モデルは最初の L トークンについてだけ答える。
4 つ目はトークン化前のスライスだ。§01 の文のバイトを任意の位置で切り、前半をデコードすると何が戻るか見てほしい:
この形式の規則に注目:バイトが文字を開始するのは、上位 2 ビットが 10 でないとき、ちょうどそのときだ。9 のうち 3 つがに落ち、ASCII では 1 つも落ちない —— 英語で試したチャンカーが、同じ文書に黙って別のトークン列を返す理由だ。