オプティマイザと訓練テクニック 入門

w ← w − η · ∇L を包むつまみは、どれも「無いと壊れる」から在る。本編が示すのは 3 つ。Adam の更新量は2 つの平均の比でスケールに依らない。最初の 1 歩は必ず ±1、それがウォームアップの理由。bf16 の重みは 3e-4 の更新を吸収できない —— マスターコピーが fp32 である理由。

01

Adam は何を平均しているのか

1 つのパラメータの勾配の移動平均 2 つと、その比。オプティマイザの中身はそれだけだ。

w ← w − η · ∇L が隠すのは、1 つの重みを 1 ミニバッチで測った ∇L の大半はノイズという事実。Adam は単一の勾配でなく、勾配列の移動平均を 2 つ持つ —— 勾配そのものと、その 2 乗。

1 つ目は指数加重平均。β₁ を残し、新しい勾配を 1 − β₁ だけ取り込むので、過去のステップは幾何級数的な割合でいまも残る。円 1 つが勾配 1 つのレジスタ内の割合。 β₁ を上げ、最新のステップが割合を返す様子を見てほしい:

β₁ = 0.00、窓は 1.0 ステップ

β₁ = 0 では円 1 つが 100% を独占し、レジスタは最新の勾配そのもの 0.769 になる点に注目。 10 ステップ窓の では最新が 13.9%、最古の 1 つも 4.4% を供給し、レジスタは 0.470。モメンタムの正体はこれだ。

2 つ目は g ではなく g² の平均で、役に立つのはその平方根。√v̂ はこのパラメータの勾配がどのくらいの大きさかを符号抜きで見積もる。 β₂ を上げ、2 つ目のレジスタが 1 つ目の隣で満ちていく様子を見てほしい:

β₂ = 0.00、窓は 1.0 ステップ

β₂ = 0 の √v̂ は最新サンプルの絶対値 0.769 そのもの、 1 回の引きしか報告しない。 では最新の取り分が 10.9%、レジスタは 0.677 に落ち着く —— 典型的な勾配の大きさが、重み 1 個あたり積和 1 回と数値 1 個で手に入る。

一方の値をもう一方で割ると単位が消える。m̂ / √v̂ は大きさ 1 程度の純粋な数で、勾配がどの桁に住もうと変わらない。勾配列ぜんぶを 50 分の 1 にし、 2 つのレジスタがそろって潰れる一方で床のゲージが動かない様子を見てほしい:

勾配はそのまま

m̂ は 4.70e−1 から 9.40e−3 へ、√v̂ は 6.77e−1 から 1.35e−2 へ —— どちらもちょうど 50 分の 1 になるのに、ゲージは 0.69 のまま。だから η は急なパラメータと緩いパラメータの妥協点でなく、全パラメータが進む 1 つの距離になる。

補正があと 1 つ残る。2 つのレジスタは空から始まるため、窓が満ちるまでまだ起きていないステップまで含めて割ってしまう ——1 − βᵗ で割ればそのずれはちょうど消える。窓を巻き戻し、補正前の総和の下に空のスロットが現れる様子を見てほしい:

ステップ 1

1 − 0.9¹ は 0.1 なので、 の補正前の総和は 0.135 で本来の 10 分の 1、補正後のレジスタはすでに 1.353 —— 勾配そのもの。ステップ 12 では 0.337、補正後 0.470 の 0.718 倍。m より v に効く —— §03 が引き継ぐ。

02

減衰をどこに置くか

更新則の中に、重みが大きくなるのを止めるものは何もない。 weight decay がその項であり、どこに置くかがその働きを決める。

勾配が同じ向きを指すかぎり Adam はほぼ同じ大きさのステップを渡し続け、m̂ / √v̂ は 重みが育っても縮まない。だから有用な重みは際限なく育ち、重みが巨大なネットワークは悪いバッチ 1 つで活性があふれる寸前にいる。

処方は、重みとともに大きくなる項。適応ステップのあとに毎回 η · λ · w を引く。勾配が何をしたかとは無関係に。λ を上げ、重みが登るのをやめる様子を見てほしい:

λ = 0.00

縮むのではなく落ち着く点に注目。押しと引きは E|m̂/√v̂| / λ で釣り合うので、 では重みがその天井のすぐ下、 1.70 付近で止まる。減衰なしは 3.51 に達しなお登る。 weight decay はゼロへの引き寄せでなく、高さ 1/λ の天井だ。

ここからは、この分野が 3 年かけて気づいた点。Adam の原論文は減衰を L2 として勾配に λ · w を足す形で畳み込んだ。その和は他のすべてと同じ 1/√v̂ の除算を通る。2 つを切り替え、1 つの λ が6 つのパラメータ群に何をするか見てほしい:

AdamW —— 分離型 · λ = 0.10

AdamW ではどの群も 1 ステップで自分の 3e-5 を失う。勾配に L2 を入れると実際の減衰は η · λ / √v̂。勾配の大きい埋め込みは 1e-3、小さいバイアスは 7.5e-1 —— 毎ステップ自分の 4 分の 3 を失う。 λ は 1 つ、ばらつきは 750 倍、エラーは出ない。

これが AdamW をどこでも既定にした修正。代償はメモリ —— パラメータ 1 個が占める 16 バイトのうち12 バイトはオプティマイザのもので、モデルのではない。モデル規模をドラッグし、柱がH100 1 枚の 80 GiB に対してどう立つか見てほしい:

7B パラメータ · 各 GPU が丸ごと持つ

どのセルが灰色かに注目。bf16 の重みと勾配の 4 バイトがモデル本体、残る 12 バイトはオプティマイザのもの。7B では 104 GiB 対 80 GiB、柱の 4 分の 1 が線を超える。 なら 1,043 GiB。

答えは複製をやめること。FSDP は 16 バイトを分割し、各 GPU は 16P/N だけを持つ。で 104 GiB は 52.2 GiB、線の下だ。請求は帯域で届く —— §07 の算数だ。

03

なぜステップ 1 が危ないのか

二次モーメントは 1 ステップ目から不偏だが、役に立たない。ウォームアップはその差を埋める傾斜だ。

§01 のバイアス補正は厳密 ——期待値の上では。代数で直せないのは、ステップ 1 で2 つ目の平均にサンプルが 1 つしかなく、 Adam がその平方根で割る点。

t = 1 では m̂ = g₁、v̂ = g₁² なので、更新量は勾配の値によらず全パラメータで±1。β₂ を変えて、冒頭のステップが動かない様子を見てほしい:

β₂ = 0.950

スライダーが届くどの β₂ でも冒頭の値は 1.00 の点に注目。まだ平均するものが無く、この比はノイズ 1 個の符号でしかない —— 傾斜が無ければ、 70 億個のパラメータが一斉にランダムな向きへ η まるごと 1 歩を踏む。最初の 40 個の更新量のうち 30 は、なお 0.5 を超える。

サンプル 1 個は極端な場合。役に立つ問いは、除数が信用できるまで何ステップ要るか。独立な 12 個のパラメータを追い、√v̂ の帯が真値に閉じる様子を見てほしい:

β₂ = 0.950 · ステップ 1

ステップ 1 では 12 本が真値の 0.01〜2.86 倍に散らばる —— 最悪は 186% ずれ、1 本は 99% 下 —— しかもどの β₂ でもそこから始まる —— ステップ 1 は各々が自分の 1 サンプルだからだ。 β₂ が効くのはその後、 では帯が β₂ = 0.999 で ±2% まで閉じ、0.95 ではまだ ±19%。短い窓は再サンプリングをやめない。

だから対処は「よい除数」でなく、除数がだめなあいだ歩幅を小さくすること。η をゼロから線形に立ち上げ、1 歩が実際に重みを動かす距離が平らになる様子を見てほしい:

ウォームアップ 0 ステップ · ステップ 1

ウォームアップ無しなら最初の 1 歩は各重みを 3.00e-4 まるごと動かす。なら 1.50e-7、ちょうど両者の比だけ小さい。実際の LLM は 500〜2,000 ステップ —— 全体の 0.1〜2%。偶然ではなく β₂ = 0.999 が平均する 1,000 ステップの窓とほぼ同じ長さだ。

04

曲線ぜんぶを選ぶ

傾斜のあとは1 本の下り坂。どの坂かということは、どこで終わりどこから始まるかほどには効かない。

ウォームアップが覆うのは最初の 1% 未満。残りの 99% は η の減衰で、どの訓練レポートも同じ 3 つの数を出す ——ピーク、形、下限だ。

使われる 4 つの形は、名前が示すほど違わない。カーソルを 10 万ステップ分歩かせ、その下で曲線を切り替えてほしい:

コサイン · ステップ 0 / 100,000

曲線の積分値に注目。コサインも線形も 16.47 —— 同じ予算の使い方が違うだけで、コサインは前半を高く保ち終盤で急落する。外れ値は WSD の 27.05。最後の 5 分の 1 までピークを保つからで、それが再開できる理由でもある。開始時に総ステップ数を知らなくてよい。

曲線がどこで終わるかは形より大きなレバー。下限をゼロから持ち上げ、曲線の下の面積が育つ様子を見てほしい:

下限 = ピークの 0%

その面積はオプティマイザが進んでよい総距離そのものなので、は 9.8% を買い足す —— すべて尾の側、モデルが探索でなく磨き込みをする区間で使われる。トークン予算が固定ならゼロまで落とし、続けるかもしれないなら下限を残す。

悩む価値があるのはピークで、自由には選べない。モデルが大きいほど下がる。GPT-3 の 8 モデルを辿り、両方の軸を読んでほしい:

モデル 4 / 8

両軸とも対数なので、直線は比例でなくべき則 —— 右へ 1 目盛でモデルは 10 倍、線は一定の割合だけ下がる。 8 点の当てはめは η ∝ N^−0.31。 から の 1,400 倍で、ピークはちょうど 10 分の 1 —— GPT-3 論文の表 2.1 に裏打ちされた経験則だ。

05

訓練を救う 1 行

勾配のクリッピングは、訓練が健全なあいだは何のコストも払わず、そうでなくなった瞬間に全部を払う。そして間違え方はちょうど 1 通りしかない。

Adam が抑えるのはパラメータごとの更新量で、更新ベクトルの全長ではない。悪いバッチ 1 つで勾配の全座標が同時に大きくなり得て、その跳躍から損失曲線は戻らない。クリッピングが抑えるのはベクトル全体だ。

規則は算術 1 行 —— ‖g‖ が閾値を超えたら、勾配全体に c/‖g‖ を掛ける。この勾配はすでに超えている。勾配を球のまわりで動かし、何が戻るか見てほしい:

‖g‖ = 2.24. 勾配の先端をドラッグ。矢印キーで 0.1 ずつ、Home キーで最初の位置に戻る
‖g‖ = 2.24 · 閾値 = 1.00 · 向きは保たれる

変わらないほうに注目。クリップされたベクトルは同じ半直線上、角度 0.0°。クリッピングが変えるのは歩幅で、向きではない。この節の残りはすべてその性質に乗る。

実際の訓練では閾値はほとんどの時間なにもしない。200 ステップ分の勾配ノルムに対して下げていき、何回捕まえるか見てほしい:

閾値 = 1.00

標準の閾値 1.00 では 200 ステップ中 4 回発火 ——スパイクの分で最大は 7.39 —— 残り 196 はそのまま。 まで下げると 200 中 180 回発火する。どのステップも同じ長さになり、実効学習率は c·η/‖g‖、訓練は静かに遅いものへ変わる —— それを告げるエラーは無い。

間違え方は 1 通りで、どのフレームワークでも識別子 1 つ隣。この図は最初から clip_grad_value_ —— 球でなく箱だ。勾配を箱のまわりでドラッグし向きが動く様子を見てほしい。ノルムに戻せば回転はゼロになる:

‖g‖ = 2.24. 勾配の先端をドラッグ。矢印キーで 0.1 ずつ、Home キーで最初の位置に戻る
座標ごとに · ‖g‖ = 2.24 · 向きが変わった

clip_grad_value_ は座標ごとに切り詰める。勾配 (2.0, 1.0) は (1.0, 1.0)、逆伝播の向きから18.4° ずれる。小さい 1 歩でなく別の 1 歩。clip_grad_norm_ を使うこと。

06

4 バイトではなく 2 バイトで

16 ビットを範囲と精度に割る。 fp16 と bf16 の違いは、その割り目から生まれる。

ここまでは数値が厳密だと仮定してきた。そうではない。現代の訓練は重み・活性・勾配を 2 バイトで持つ —— メモリ半分、同じ演算器の上でスループットはおよそ倍。H100 SXM の tensor core は bf16 で 989 TFLOPS、TF32 で 495。ベクトル fp32 の 67 TFLOPS は別の通路で、その 15 倍はどの行列積も見ない。

浮動小数点は指数ビットで範囲を、仮数ビットで精度を買う。16 ビットしかなく、片方が得れば他方が失う。割り目を動かし、届く区間が伸びる様子を見てほしい:

指数 5 ビット、仮数 10 ビット

fp16 と bf16 はこの 1 本の線の上の 2 点。fp16 は と仮数 10 ビット、bf16 は で、指数部は fp32 と同じ。だから bf16 は 3.4e38 まで届き、fp16 は 6.6e4 で止まる。代わりに十進の有効桁は 2.4、fp16 の 3.3 に及ばない。

精度は抽象概念でなく、その形式が持てる隣り合う 2 値の間隔そのもの。軸に沿って数をドラッグし、変換がそれをどこへ置くか見てほしい:

x = 1.288e+0. 軸に沿ってドラッグ。矢印キーで 10 分の 1 桁ずつ、Home キーで最初の値に戻る
x = 1.288e+0 · bf16

この間隔は値に対する一定の割合なので —— bf16 ではどこでも 2⁻⁷、つまり 0.78% —— 相対誤差は 1e-30 でも 1e30 でも同じで、 fp16 の 0.098% の 8 倍。どの数も 2 本の目盛のどちらかに着地する。あいだには何も着地しない。

実務で fp16 を壊すのは範囲。勾配は 1 より数桁下に住み、fp16 最小の非正規数は 5.96e-8 しかない。ロススケールを動かし、それが分布を2 枚の壁のあいだで運ぶ様子を見てほしい:

損失スケール = 2^0 · fp16

スケーリング無しでは、このモデル分布の 8.0% が fp16 の床を下回る —— 静かにゼロへ落ち、そのパラメータは訓練されない。逆伝播の前に損失を 倍すれば 0.05%、 まで押すと上側で 5.9% が inf にあふれる。bf16 はどの設定でも分布ぜんぶを保つ。勝った理由のほとんどがこれだ。

どの形式でも fp32 に留めるものが 1 つある。その理由がこの図。1.0 という重みに更新量を足したとき、実際の増分を比べてほしい:

更新量 = 1e-3 · fp32 マスター

1e-3 の更新は fp16 で 9.77e-4、bf16 ではまったく着地しない —— 0.0078 の間隔の半分に満たず、加算は重みをそのまま返す。η が 3e-4 なら実際の更新の大半がこの大きさ。だから作業用は bf16、マスターコピーは fp32 にする。

07

1 つのモデルを切る 3 つの切り方

データ、テンソル、パイプライン。独立な 3 つの切り口と、主に線上のバイト数の算数が、どれをどこで切るか決める。

§02 は 7B の 104 GiB を H100 の 80 GiB に対して数え、データ並列群への分割で答えた。メモリを帯域で買う取引が、この節の通貨だ。訓練は 3 本の軸で切られ、各軸の大きさは線の上で何を言うかが決める。

いちばん簡単なのはデータ並列。全 GPU がモデル全体を持ち、バッチの一部を回して勾配を平均する。面白いのはこの「平均」——リングを 1 歩ずつ進め、各チャンクが 1 ホップごとに寄与を 1 つ拾う様子を見てほしい:

GPU 4 枚 · フェーズ 0

ハブが居ない点に注目。最初の N−1 フェーズで各チャンクはリングを 1 歩ずつ回り、止まるたびに GPU 1 枚分の寄与を拾う。フェーズ 3 ではどの GPU も完全に縮約されたチャンクを 1 つ持つ。続く N−1 フェーズはその完成品を回して配る。調整役を通るものは何もない。

これが効くのは代わりの方式の代償ゆえだ。GPU 枚数を増やし、リングと中央サーバが両対数のフレーム上で離れていく様子を見てほしい:

GPU 8 枚

リングで最も混むリンクが運ぶのは 2(N−1)/N · M、2M に向けて登りそこで止まる。bf16 勾配 13 GiB は GPU 8 枚でリンクあたり 22.8 GiB、でも 26.0 GiB。中央サーバは 2NM —— 8 枚で 209 GiB、512 枚で 13.0 TiB。片方は N で平ら、他方は両対数で上りの直線 —— all-reduce がリングである理由だ。

データ並列はモデルが GPU 1 枚に載ることを要求し、載らなければ重み行列を切る。各 GPU がどの層でも1 枚の板を持つ形だ。層を分割し、その代償を見てほしい:

1 方向テンソル並列

隠れ次元 4096 の層は 201M パラメータ。で GPU あたり 25.2M。得たぶんは通信量で払う —— 1 層 1 イテレーションに活性の all-reduce は 2 回でなく4 回。 Megatron は各ブロックを f/g の対に分け、層にその 2 対があるので順伝播 2 回、逆伝播 2 回になる。 8 分割・8,192 トークンの 1 マイクロバッチでは 448 MiB、活性そのものは 64 MiB。どの層でもどのイテレーションでも —— テンソル並列が NVLink ノード内に留まる理由だ。

3 つ目は層に沿って切る。段 0 が先頭のブロック群、段 1 が次を持つ。この代償には通信が入らない。マイクロバッチを増やし、空の枠が閉じていく様子を見てほしい:

4 段 · マイクロバッチ 1 個

マイクロバッチが 1 つだと、4 段は75% の時間を遊ぶ。前の段が終わるまで次は始められないからだ。バブルはちょうど (p−1)/(m+p−1)。戦い方はマイクロバッチを増やすこと ——で 4 段は 16% になる。実際は 3 つを同時に使う。1,024 GPU ならデータ 32 × テンソル 8 × パイプライン 4、この積が GPU 枚数だ。

08

9 行と、それぞれが何のためにあるか

このページのつまみはすべて、1 訓練ステップのどこかに着地する。その場所を並べる。

訓練スクリプトの仕掛けは平凡だ。コンストラクタが 1 つ、ループが 1 つ、中に7 つの呼び出し。「動く訓練」にしているのは、どの呼び出しも特定の失敗に答え、その答えが正しい順に届くことだ。

1 反復を辿り、各操作が握るつまみと、パラメータあたり何バイトをメモリで動かすかを読んでほしい:

操作 1 / 7 · 順伝播

どの行が長いかに注目。1 反復 56 バイトのうちオプティマイザのステップが 28、順伝播は 2 —— Adam のステップが帯域律速である理由がこれだ。順序も恣意的ではない。スケール解除は all-reduce のあと、クリップの前。 1.0 と比べるノルムは本物でなければならない。

opt = AdamW(p, lr=3e-4, betas=(.9, .95),
            eps=1e-8, weight_decay=0.1)

for step in range(1, TOTAL + 1):
    set_lr(opt, lr_at(step))     # §03, §04
    with autocast(dtype=bfloat16):  # §06
        loss = model(batch).loss
    loss.backward()              # §07 ring
    clip_grad_norm_(p, 1.0)      # §05
    opt.step(); opt.zero_grad()  # §01, §02

研究所が公表し、ファインチューニングが受け継ぐ数だ。