勾配降下 入門
1 行の算術を数十万回繰り返すことで、耳にしたことのあるモデルはすべて訓練されている。この入門ではそれを分解する:1 歩そのもの、収束するか爆発するかを決める 1 つの数、モーメンタムと Adam が存在する理由、1 歩あたりのデータ量、そして走行を終わりまで運ぶスケジュールだ。
下りへ 1 歩
耳にしたことのあるモデルはすべて、1 行の算術の繰り返しで訓練されている:傾きを測り、その逆へ 1 歩進み、また測る。
斜面に乗ったビー玉を思い浮かべる。ビー玉はパラメータベクトル θ、斜面は損失 L(θ) で θ ごとに数が 1 つ、「下り」は勾配 ∇L の反対方向。1 次元なら斜面は曲線で、ビー玉はその上の点だ。
ビー玉の足元の傾きが、このアルゴリズムの唯一の入力だ。接線は両脇で急、最小点で水平になる。ボールをドラッグして、その下の接線から ∇L を読む:
∇L は符号つきの数であって方向を表す言葉ではない:θ = 3.20 では +3.20、左側では負になる。勾配は常に上りを指す。そしてボールが最小点に近づくほど短くなる —— 更新規則が最初に利用するのが、この性質だ。
規則は θ ← θ − η∇L:勾配に学習率 η を掛けて引く。η = 0 では何も動かない。上げていって歩幅が伸びるのを見る:
正の勾配を左向きの矢印に変えているのはマイナス符号だ。着地点を見ると、η = 1.00 でちょうど最小点に乗り、1.05 では通り過ぎる。歩幅は積 —— こちらが選ぶつまみと、測った傾きの積 —— なので、同じ η でも問題が変われば振る舞いは変わる。
すべてがこの符号に懸かっている。スイッチは同じ 4 歩をθ ← θ − η∇L とθ ← θ + η∇L で走らせる:
プラス符号は派手に失敗しない。走るし、数も返す。そして損失は 4 歩で 5.12 から 15.66 まで登る —— 損失が上がっているのにエラーは 1 つも出ない訓練だ。フレームワークは符号をオプティマイザの中に隠すので、実際に出会うのは自分で書いた勾配の符号が逆になっている版になる。
下る側に戻ると、もう 1 つ性質がただで付いてくる。η = 0.35 固定で 12 歩を 1 歩ずつ進め、この 1 歩の長さを読んでいく:
勾配が最小点までの距離とともに縮むので、歩幅も縮む —— 1 歩目 1.12、12 歩目 0.010。誰かが仕組んだのではなく、規則自身の算術だ。走行が収束済みに見えるのが実際よりずっと早いのも、これが理由だ。
不変条件はそのまま表明になる:曲率が高々 a の損失なら、η < 2/a のどの 1 歩もパラメータを最小点へ近づける ——L(θ − η∇L) − L(θ) ≤ −η(1 − ηa/2)‖∇L‖²。この評価には a が要るが、誰も渡してくれない。
1 歩の大きさ
訓練が収束するか、這うか、爆発するかを 1 つの数が決める —— しかも同じ数が、損失が違えば違う意味になる。
§01 のボウルは曲率 a = 1、更新はちょうど θ ← (1 − η)θ だ。この係数がすべてで、最小点までの距離は毎ステップ |1 − η| 倍される。
同じボウルで 12 歩、学習率は手の中。小さければ這うだけ、 なら 1 歩、 を越えるとボールはボウルの外へ出る:
1.00 と 2.00 のあいだで何が起きるかを見る:どの歩も最小点を通り越して反対側に着地するが、それでも前より近い。行き過ぎることは発散ではない。左右に振れながら縮んでいく走行は収束していて、振れながら育つ走行はしていない —— その境界は厳密だ。
その境界は描ける数でもある。|1 − ηa| が距離に毎ステップ掛かる係数だ —— マーカーをその上でドラッグする:
V の底は η = 1/a にあり、そこでは係数が 0 で 1 歩で厳密に解ける。その両側では進みは幾何的だ:η = 0.20 なら係数 0.80 で 1% まで 21 ステップ、η = 0.05 なら 0.95 で同じ 1% に 90 ステップ。学習率を 4 分の 1 にすると、ステップ数は 4 倍になる。
訓練中に見えているのはそれではない。見えるのは損失 —— 1 ステップに 1 つの数を、対数軸に並べたものだ。η を決めて形を読む:
形は 3 つ、つまみは 1 つ。η = 1 より下ではまっすぐ落ちる。対数軸では幾何的減衰が直線だからだ。1 と 2 のあいだでは鋸歯を描いて下がる。2 を越えると登り、しかも幾何的に登る ——発散した走行が数千ではなく数十ステップで inf に達するのはそのためだ。
境界は 2/a なので、安全な学習率は曲面の性質であって、オプティマイザの性質ではない。曲率と学習率の面で点をドラッグし、境界を越えてみる:
2 本の曲線は同じ形だ。η = 2/a が発散する場所、η = 1/a が 1 歩で解ける場所で、どちらも曲率とともに動く。曲率が倍なら学習率は半分 —— あるモデルの学習率が規模の違うモデルに移せない理由だ。
実際の a は 1 つの数ではなく、走行が訪れる範囲での最大の曲率で、走行とともに変わる。だから典型的な失敗はステップ 0 での発散ではなく、1 時間訓練してからより尖った場所で爆発することだ。
2 つの方向、1 つの学習率
本物のモデルには数百万のパラメータがあり、曲率は共有していない。1 つの学習率がその全部を一度に受け持つ。
問題の全体は 2 次元で見える。損失は ½(axx² + ayy²) —— 一方向に長く、もう一方を横切る向きに急な谷だ。比 κ = ay/ax が条件数になる。
まず幾何から。点を谷の好きな場所へドラッグして、勾配がどちらへ送ろうとするかを見る:
矢印は谷に沿ってではなく、ほぼ真横に谷を横切る向きを指す。開始姿勢では横向きが縦向きの 8.5 倍長く、つまり 1 歩はほぼ垂直 —— すでにほぼ正しい方向に費やされている。勾配降下は最小点を指さない。最も急な壁の下を指す。
その点から 40 歩、学習率は手の中。急な方向がどこまで大きくできるかを決める —— を越えると発散する:
急な方向が天井を決めるので、この谷が許す最良の学習率は で、それでも最後の 1% を詰めるのに 28 ステップかかる。緩い方向が、急な方向のために選ばれた学習率で這わされているからだ。η を 0.40 まで下げると鋸歯は消えるが、ステップ数は 130 になる。
どこまで悪くなるかは 1 つの数の関数だ。κ ごとに最良の学習率をとり、谷を円から まで引き伸ばす:
ステップ数が κ とともに増えていくのを見てほしい。円は 1 歩で解ける —— 曲率が 1 つなら完全な学習率も 1 つ —— そこから:κ = 2 で 5 歩、12 で 28 歩、40 で 93 歩。1 歩あたりの収縮が (κ − 1)/(κ + 1) で、下から 1 に近づいていくからだ。最適化を遅くするのは次元ではなく条件数だ。
安い変更 1 つで大半は直る。前の 1 歩の β 分を残して今の 1 歩に足す —— これが速度だ —— そして生の勾配ではなくそれに沿って進む:
横切る成分が打ち消し合うのを見てほしい。連続する 2 歩は同じ壁を上下に指すので引き算になり、谷に沿う成分は同じ向きなので速度が積み上がる。 では同じ走行が 28 歩ではなく 9 歩で済む —— 鋸歯が自分自身ではなく前進に使われている。
モーメンタムには固有の失敗があり、それは発散ではない。β をこの谷が望む値より押し上げて、最小点までの距離を対数軸で読む:
β = 0.90 でも収束はする —— 9 歩ではなく 74 歩で、大半を最小点から離れては戻ることに使う。過剰なモーメンタムはボウルの中の重い球だ。兆候は数十ステップで振動する損失曲線。
理論では、正しい組は κ を √κ に変える:η = 4/(√amax + √amin)² と β = ((√κ−1)/(√κ+1))² で収縮は (√κ−1)/(√κ+1)。κ = 12 なら 28 歩に対し 12 歩。難点はどちらにも κ が要ることだ。
パラメータごとの学習率
本物のネットワークの条件数は巨大で、最悪なのはパラメータごとに必要な歩幅がまるで違うことだ。
1 つのモデルからパラメータを 6 つ取る —— LayerNorm のゲイン、バイアス、アテンションの射影、稀な token の埋め込み行。同じ逆伝播でそれらの勾配は 2 桁に散らばり、1 つの η が 6 つ全部に掛かる。
どれもが 0.05 前後の歩幅を 2 倍以内で欲しがるとする。 を動かし、6 つのうち何個が帯に入るか数える:
η を動かしながら数を見てほしい:よくても 6 つ中 2 つだ。勾配が 120 倍に散らばれば歩幅も 120 倍に散らばり、帯の幅は 4 倍しかない —— η をどう選んでもこの算術には勝てない。いちばん大きいパラメータが行き過ぎる歩幅を取るか、いちばん小さいものが終わらない歩幅を取るかだ。
Adam の答えは、勾配自身の大きさで割ることだ。パラメータごとに移動平均 m と移動二乗平均 v を持つ。カウンタをスクラブして、√v が|g| の上まで歩いていくのを見る:
√v が勾配自身の大きさに収束するので、m̂/√v̂ はスケールによらず ±1 になり、どのパラメータも 1 ステップに学習率 1 つ分だけ動く。Adam が買っているのはこれだ —— 速度ではなく、自分では選んでいないスケールに依存しない歩幅。η は「勾配の何倍」ではなく「1 ステップあたり」を意味しはじめる。
累積器はどちらも 0 から始まるので、序盤は小さめに出る —— しかも 2 つの減衰の速さは同じではない。補正を切ったまま、カウンタを最初の 200 ステップに通す:
β₂ = 0.999 は β₁ = 0.9 より 100 倍ゆっくり立ち上がるので、小さすぎるのは分母のほうだ。補正なしでは Adam の 1 歩目は学習率 3.16 個分で、 でピークになる。失敗は静かだ:訓練は走り、損失は下がり、最初の数百ステップは設定の 6 倍の速さで踏まれている。
この状態はただではない。累積器はどちらもパラメータの完全なコピーで、しかも fp32 だ。モデルとオプティマイザを選んで、80 GB のカード 1 枚に対する請求書を読む:
Adam はパラメータあたり 16 バイト、SGD は 8 —— 重みと勾配の fp16、 fp32 のマスターコピー、fp32 のモーメント 2 つだ。 なら Adam に100 GB、SGD に 50 GB —— この節の主題を手放せば収まる。
つまずきどころもここだ。Adam の weight decay は L2 正則化ではない:λθ を勾配に足すと同じ 1/√v̂ の割り算を通るので、勾配の大きいパラメータほど減衰が弱くなる。AdamW は割り算のあとで λθ を引く —— 1 行の差が精度に効く。
1 歩あたりのデータ量
更新規則の勾配はデータセット全体の平均だ。誰もそれを計算しない。実際のステップはすべて標本を使う。
パラメータ 1 つ、サンプル 64 個、サンプルごとに勾配 1 つ。全バッチ勾配はその平均 ——規則が本当に欲しがる数だ。ミニバッチは B 個を取り、その平均で代用する。
バッチを大きくして、推定値が真の値へ歩み寄るのを見る:
B = 1 では推定値は 0.452、真の値は 0.528 —— しかもこの誤差の符号は引くたびに変わる。SGD が働くのは、誤差の平均が 0 だからだ。推定値は不偏なので、多数のステップにわたって間違いは打ち消し合い、正しい成分だけが積み上がる。ノイズはバイアスではない。
誤差が縮む速さはふつうの平方根則だ。マーカーを直線に沿ってドラッグし、標準誤差を読む:
両軸とも対数なので、このべき乗則は傾き −½ の直線になる:が買うのはノイズ半分で、 4 分の 1 ではない。バッチを倍にすると計算は倍かかり精度は 1.41 倍、この比は規模が変わっても良くならない。
語彙も同じ絵から出てくる。データを 1 周ぶんバッチに切る: 1 バッチが 1 イテレーションであり 1 回の更新、行全体が1 エポックだ:
1 エポックはデータを 1 周すること、1 ステップは 1 回の更新。単位が違い、両者の比は ⌈N/B⌉ になる —— だから「10 エポック訓練した」は、B を知らないかぎり何回更新が起きたのかを何も言っていない。スケジュールはステップで索引され、エポックでは索引されない。
実際の数を入れてみる。ImageNet の訓練分割は 1,281,167 枚で、古典的なレシピは 90 エポック回す ——バッチサイズをドラッグして両方の数を読む:
B = 256 なら 1 エポックあたり 5,005 ステップ、走行全体で450,450 回の更新。バッチを にすると同じ 90 エポックが 28,170 回の更新になる —— まったく同じデータの上で、オプティマイザが動ける機会が 16 分の 1 だ。
それには理由があって買われている。1 ステップには B に依存しない固定費がある —— カーネル起動、オプティマイザ自身の要素ごとの処理、勾配の all-reduce 1 回 —— そして B に依存する算術があり、それがスループットを上限へ押し上げる:
曲線がどこで折れるかに注目したい。B = 80 より下では固定の 12 ms が支配し、スループットはほぼバッチに比例する:32 → 256 で 2.7 倍。それより上は算術が支配し、曲線は上限に張り付く: はバッチ 4 倍で 1.22 倍、更新回数は 4 分の 1 だ。
バッチサイズは 3 者間の取引だ:ノイズは 1/√B で下がり、 1 エポックの実時間は飽和まで下がり、更新回数は 1/B で下がる。通常の解決は、バッチとともに学習率を線形に上げることだ。
終わりまで持っていく
収束する走行と終わる走行は別物だ。いつ何の値で止まるかを 3 つが決める。
ここまで η を走行中に変えていない。大きなモデルはどれも変える:ほぼ 0 から数千ステップかけて登り、最後はほとんど 0 まで減衰させる。前半と後半は別々の問題を直している。
まずウォームアップの側。ウォームアップの長さをドラッグして、最初の 1 歩がどこまで許されるかを見る:
ウォームアップがなければ、1 歩目はピークの学習率で踏まれる —— ランダム初期化の中へ、1 バッチから推定した勾配で、Adam の二次モーメントがまだ空のままで。最後の 1 つだけで 3〜6 倍ぶんの価値があることは §04 が示した。ウォームアップは迷信ではなく、最初の数百ステップを生き延びるいちばん安い方法だ。
減衰の側が直しているのは別のことだ。一定の学習率でノイズありの 200 ステップを走らせ、損失がどこで下がり止まるかを見る:
損失は最小点に収束しない。最小点のまわりの雲に収束し、その雲の大きさを学習率が決める。床は ησ²/2(2 − ηa) ——床も半分になる。減衰スケジュールが買っているのはこれで、走行の最後の 5 分の 1 がほぼ学習率なしでもまだ数字を動かせる理由でもある。
残りは 1 つ、そしてそれが人を叩き起こすやつだ。ステップ 24 に、いつもの 26 倍の勾配を持つ悪いバッチが来る。クリップ閾値を決める:
クリップしなければ、そのバッチ 1 つが 2.00 の歩幅を踏み、損失を 0.004 から1.84 へ投げ返す —— サンプル 1 個で数時間ぶんの前進が消える。クリッピングは閾値より長い勾配を閾値まで縮めるので、 は η·c になる。
では収束したとどう分かるのか。損失では分からない —— 損失はノイズの床に支配されている。正直な信号は勾配ノルムで、本当の最小点ではそれが 0 に向かうのに損失は向かわない。
ループと、3 つのオプティマイザ
置き場所を間違えた 2 行、1 つの問題を 3 つのオプティマイザで、そしてあらゆる訓練ループである 8 行。
下の 2 つの間違いはどちらも 1 行の置き場所だけで、どちらもエラーなしに走る。スイッチが片方を元の場所に戻す:
zero_grad() の欠落は派手だ —— 勾配が累積して実効学習率が育ち、損失は 2.0 に居座る。エポックごとの減衰は静かで、収束はする。ただし 6.1e−4 ではなく8.7e−3 —— §06 の床から離れないからだ。
下の各行は上のどれかの節だ:
for epoch in range(epochs):
for x, y in loader: # 1 iteration
opt.zero_grad() # grads add up
loss = criterion(model(x), y)
loss.backward() # fills p.grad
clip_grad_norm_(params, 1.0)
opt.step() # theta -= lr*g
sched.step() # per stepあと 2 つ:Adam の weight decay は AdamW ではない。バッチ 256 の学習率は 2,048 のそれでもない。最後に §03 の谷を、3 つそれぞれの最良の定数学習率で:
モーメンタムの完勝だ:SGD の 28 歩、Adam の 98 歩に対して 12 歩。Adam は正規化された手法で、勾配がどうであれ各座標を 1 ステップにつき η 程度動かす。だからステップ数は条件数ではなく距離に比例する。勝ち筋は §04 のほうだ:テンソルごとの調整が要らない。