コンテンツへスキップ
6/30第6章 / 全30章

学習させること、そして汎化させること

損失が ln 2 から動かない6層ネットワークを、測定ごとに修正。続いて double descent:40点に5,000 parameters。

このページの内容

Chapter 5 のネットワークは動きます。parameters は9個、XOR を学習し、gradients は PyTorch と小数点以下16桁まで一致します。

それを6層深くすると、学習が完全に止まります。遅いのではなく、完全にです。これは2スパイラル分類問題に対する6層ネットワークを、5000 steps 学習させたものです。

TEXT
step    1: loss 0.693147
step 5000: loss 0.693147
accuracy: 50.0 %

この数値は任意ではありません。ln2=0.693147\ln 2 = 0.693147 は、すべてに確率 0.50.5 を出力するモデルの二値クロスエントロピーであり、50 % はバランスしたデータセットでのコイントスです。5千 steps 後でも、ネットワークは1桁も動いていません。クラッシュも警告もなく、gradients はなお完全に正しいままです。

この章は、走るネットワークと、機能するネットワークの差についてです。別々のテーマに見えて同じ仕事である2つの半分があります。損失を下げること、そしてモデルが一度も見たことのないデータでもそれを下げることです。

推測するのではなく、まず見ます。入力のバッチを流し、各層の活性化の標準偏差を出力し、次に重み gradients の標準偏差を出力します。

profile.pyPYTHON
def profile(model, x):
    h = x
    for layer in model:
        h = layer(h)
        if isinstance(layer, (nn.Tanh, nn.ReLU)):
            print(f"activation std: {h.std().item():.4f}")
    model(x).sum().backward()
    for p in model.parameters():
        if p.dim() == 2:
            print(f"gradient std: {p.grad.std().item():.2e}")

3つの初期化、同じアーキテクチャ、tanh\tanh の6層です。

初期化活性化 std、層 1→6
normal, std 0.010.010.0145 · 0.0016 · 0.0002 · 0.0000 · 0.0000 · 0.0000
normal, std 110.6573 · 0.9296 · 0.9585 · 0.9634 · 0.9637 · 0.9625
Xavier0.1579 · 0.1493 · 0.1353 · 0.1333 · 0.1325 · 0.1403
初期化gradient std、最初の層 → 最後
normal, std 0.010.013.20e-06 · 4.97e-07 · … · 6.40e-06
normal, std 111.94e+03 · 2.28e+02 · 1.22e+02 · 4.43e+01 · 1.85e+01 · 7.30e+00
Xavier2.31e+00 · 4.50e-01 · 4.26e-01 · 3.89e-01 · 4.39e-01 · 4.73e-01

最初の行が上のネットワークで、これはゆっくり学習しているのではありません。信号がもう残っていないのです。第4層までに、活性化の標準偏差は小数4桁表示でゼロにアンダーフローしています。どの入力も同じ出力を生み、出力は定数であり、定数の gradient は何もありません。重みは「安全のため」に小さく初期化されましたが、その小ささが致命的でした。

2行目は逆方向の失敗で、直感に反するので理解する価値があります。活性化は健全に見えます。およそ0.96です。しかしそれは tanh\tanh飽和し、限界付近に張り付いている状態であり、まさに Chapter 5 が gradient をほぼ1万分の1失う領域として測定したものです。それでも gradients は巨大です。最初の層で1940です。この2つは同時に真です。backward の各 step は WW^\top を掛け、分散1の入力128個ではその因子が約 12811\sqrt{128} \approx 11 のゲインを持つため、飽和した tanh\tanh による縮小を圧倒します。gradients は戻る途中で幾何級数的に増えます。これが exploding gradient であり、実際の学習 run では数 steps のうちに nan という損失値を生みます。

3行目が望ましい状態です。活性化のスケールは深さ方向でおおむね一定、gradients のスケールも深さ方向でおおむね一定です。何も死なず、何も爆発しません。

うまく初期化すると、step zero のスケールは直ります。しかし固定されたままにはなりません。重みは動き、5千 steps までには慎重な分散の議論はもう当てはまりません。

正規化層はスケールを継続的に強制します。活性化ベクトルが与えられたら、平均を引き、標準偏差で割り、それから学習されるスケール γ\gamma とシフト β\beta を適用します。必要なら層が正規化を打ち消せるようにするためです。

h^=hμσ2+ϵ,y=γh^+β\hat{h} = \frac{h - \mu}{\sqrt{\sigma^2 + \epsilon}}, \qquad y = \gamma\hat{h} + \beta

本当に重要な問いは、何に対して平均を取るかだけです。Batch normalisation3 はバッチ次元に沿って μ\muσ\sigma を取り、feature ごとに1つの統計を持ちます。Layer normalisation4 は features に沿って取り、例ごとに1つの統計を持ちます。

その選択は小さく見えますが、その後のほぼすべてを決めます。

BatchNorm では、各例の出力が、たまたま同じバッチに入った他の例に依存します。学習時にはこれは軽い正則化になります。推論時にはバッチがないので、学習中に集めた統計の移動平均を保持する必要があります。つまり、その層は training mode と evaluation mode で異なる振る舞いをし、mode の切り替え忘れは現場で最もよくあるバグの1つです。また小さいバッチでは劣化し、可変長 sequence では扱いにくくなります。「位置40でのバッチ平均」は、たまたまその長さまである sequence がいくつあるかで計算されるからです。

LayerNorm は各例をそれ単独で正規化します。バッチ依存なし、running statistics なし、学習と推論で同一の振る舞い、バッチサイズに無関心、sequence 長に無関心です。1人のユーザーに対して token を1つずつ生成する段階になると、これらの性質はどれも「あればよい」ではなく要件になります。それが Chapter 13 の行き着く先です。

だからこそ LayerNorm は Chapter 9そのまま再登場します。transformer ブロックがそれを使い、しかも抽象的によりよく動くからではなく、右の列にある理由で使うのです。

死んだネットワークに対する候補の修正は4つあります。Xavier initialisation、LayerNorm、residual connections、そして SGD の代わりに Adam です。誘惑されるのは、4つ全部を適用して先に進むことです。それをすると、どれが効いたのか永遠に分かりません。そして次に同じことが起きたとき、方法はなく、儀式だけが残ります。

だから1つずつ適用します。同じ seed、同じデータ、同じアーキテクチャ、800 steps です。

追加したもの最終損失accuracy
なし0.693150.0 %
Xavier initialisation0.569260.4 %
LayerNorm0.623061.5 %
residual connections0.665156.6 %
Adam0.678758.7 %
4つすべて0.0000100.0 %

午前2時に読むようにこの表を読むと、結論はこうなります。単独では何も効かず、全部一緒なら効く。したがって deep learning は錬金術である。この結論は間違っています。そしてその理由を突き止めることが、この章で最も役に立つことです。

各 run に6倍の予算、つまり800ではなく5000 steps を与えると、状況は完全に変わります。

追加したものfinal loss @ 5000accuracy
なし0.693150.0 %
Xavier initialisation0.0007100.0 %
LayerNorm0.0002100.0 %
residual connections0.665356.7 %
Adam0.690853.4 %
Xavier + Adam0.0000100.0 %
Xavier + LayerNorm0.0001100.0 %

今度は像がくっきりします。これは儀式ではなく診断です。

初期化だけで直ります。正規化だけでも直ります。 どちらも本当の病気、つまり forward 信号がゼロへ崩壊することに対処しており、どちらか一方で十分です。800 steps では部分点に見えただけで、問題はすでに解決していて、まだ抜け出している途中だったのです。

Residual connections と Adam は、どの予算でも直しません。 悪いからではなく、別の病気を治療しているからです。residual connection は、gradient に詰まった層を回避する経路を与えます。gradient が問題なら非常に価値がありますが、forward 信号がすでにゼロなら価値はありません。死んだ層を迂回する近道も、死んだ値を運ぶだけだからです。Adam は各 parameter の step を、その parameter 自身の gradient 履歴でリスケールします。gradients の大きさが大きく異なるときには役立ちますが、出力が入力に依存していないネットワークを蘇生することはできません。

そして「なし」は5千 steps 後でもまだ正確に 0.6931 です。0.6929 ではありません。遅いのではなく、死んでいます。その区別は以前より見えるようになっています。比較対象として、修正が効くことを示す行があるからです。

ここから先、このコースでは PyTorch を使います。それは宣言ではなく、獲得されるべきものです。そこで、PyTorch が何をしているのか、あなたがすでにできることに対応させて正確に示します。

optimizer とは、gradients を parameter updates に変える規則です。素朴な gradient descent は gradient を使います。Momentum はその running average を使い、ノイズをならし、一貫している方向に速度を積み上げます。

optim_by_hand.pyPYTHON
v = beta * v + p.grad          
p -= lr * v                    

Adam5 は2つの running averages、つまり gradient と gradient の二乗を保持し、一方を他方の平方根で割ります。そのため各 parameter は、自分自身の最近の gradient magnitude に合わせてスケールされた step を得ます。

optim_by_hand.pyPYTHON
m = b1 * m + (1 - b1) * g          # mean of the gradient          
v = b2 * v + (1 - b2) * g * g      # mean of the squared gradient  
m_hat = m / (1 - b1 ** t)          # bias correction: both averages start at zero
v_hat = v / (1 - b2 ** t)
p -= lr * m_hat / (v_hat.sqrt() + eps)   

10行です。同じ問題で torch.optim と照合し、50 steps 実行します。

TEXT
SGD+momentum   by hand [2.7781870365142822, -1.0304985046386719]
               torch   [2.7781870365142822, -1.0304983854293823]   max |diff| = 1.19e-07
Adam           by hand [0.4893140196800232, -0.46317872405052185]
               torch   [0.48931416869163513, -0.46317875385284424]   max |diff| = 1.49e-07

float32 精度まで同一です。torch.optim.Adam はこの5行に、edge cases への数十年分の配慮と C++ kernel を足したものです。ここから先であなたが行う交換はこれです。理解を魔法と交換するのではなく、すでに書いた行を速度と交換するのです。

Adam の通常の説明は「parameter ごとの adaptive learning rates」ですが、それは理由ではなく説明です。理由は幾何であり、測定できます。

方向によって曲率が異なる損失を考えます。一方は急で、もう一方は浅い。SGD には1つの global learning rate しかないため、最も急な方向で安定するほど小さい値を選ばなければなりません。そしてその値は浅い方向には小さすぎ、進捗は這うように遅くなります。これが gradient descent が狭い谷をジグザグに下る古典的な図を生む原因です。

2つの曲率比、3つの optimizers、300 steps。そして不利にならないよう、各 optimizer には sweep で得た最良の learning rate を与えます。

曲率比SGDSGD + momentumAdam
10 : 1error 0.000002error 0.000000error 0.000000
1000 : 1error 1.925485error 0.001432error 0.000000
diverged at (1000:1)4 of 8 rates4 of 8 rates0 of 6 rates

比が10なら、すべてうまくいき、議論することはありません。1000 では、素朴な SGD は試したどの learning rate でも答えに到達できません。最良でも error は1.93のままで、半分の rates では outright に diverge します。Adam は target にぴったり到達し、どれでも diverge しません。

最後の列こそ、Adam が default である実務上の理由です。Adam がより良い解を見つけるからではありません。条件のよい問題では、tuned SGD がしばしば同等か上回ります。Adam は、あなたが選んだ learning rate への感度がはるかに低いのです。そして実際のネットワークは、何百万もの parameters 全体にわたり、1000 をはるかに超える曲率比を持ちます。

ここに属する要素があと2つあり、どちらも1行です。Gradient clipping は、gradient vector の norm が threshold を超えるたびにリスケールし、診断表の「損失が突然巨大な値へ跳ねる」行を何事もなかったことにします。そして learning rate schedules です。最初の数百 steps では、ほぼゼロからの短い warmup を使います。Adam の分散推定は、ある程度 gradients を見るまではゴミであり、そのゴミに対して full-size step を踏むと初期化を壊しかねないからです。その後はゼロへ向かう cosine decay です。run の終わりに開始時と同じ step size を使うということは、最小値に落ち着くのではなく、その周りで震え続けることだからです。

ここまではすべて、損失を下げることについてでした。ここからはより難しい半分です。損失を下げることは目標ではなく、目標の proxy だからです。そして proxy は、特定の有名な形で失敗します。

少しノイズのある滑らかな関数から12点を取ります。次数を増やしながら多項式を fit します。

次数train RMSEtest RMSE
10.7644990.6985
30.2526050.3031
50.1644370.1568
90.0889600.2347
110.0000001.2094

12点に対する11次多項式は、すべての点を1つ残らず正確に通ります。train error は小数6桁までゼロです。しかし見たことのないデータでは、5次の8倍悪くなります。次数3と次数11に、training range のすぐ外側である x=3.25x = 3.25 を予測させます。

TEXT
degree  3: predicts   -1.053   (truth -0.012)
degree 11: predicts  +61.224   (truth -0.012)

答えがおよそゼロであるところに、61です。モデルは関数を学習したのではありません。12点を学習し、その間では算術が要求することを何でもしているだけです。

これが overfitting です。そしてその反対、曲線をまったく表現できずどこでも悪い次数1が underfitting です。古典的な説明では、モデルの期待 error を3つに分けます。bias は、モデルが真実を表現するには硬すぎることによる error。variance は、モデルが柔軟すぎ、この特定のサンプルのノイズを追いかけることによる error。そして irreducible noise は、何をしても直らないものです。単純なモデルは bias が高く、柔軟なモデルは high-variance で、古典的な処方は真ん中の sweet spot、上の表でいう次数5を見つけることです。

標準的な道具は、どれも variance 項を攻撃します。

  • L2 正則化(weight decay)は損失に λw2\lambda \lVert w \rVert^2 を足し、weights をゼロへ引っ張って関数を滑らかにします。上の表では、次数11の最大係数が損害を与えています。大きさを penalize すれば、それを無力化できます。
  • L1 は代わりに λwi\lambda \sum |w_i| を足します。違いは見た目だけではありません。L2 の gradient は weight に比例するため、weight が小さくなるにつれて縮み、ゼロに近づきますが到達はしません。一方 L1 の gradient は一定の ±λ\pm\lambda で、最後まで押し続けます。そのため L1 は weights を正確にゼロにします。feature を選択するのです。L2 は小さい weights を生みます。滑らかさが欲しいなら L2、疎性が欲しいなら L1 を使います。
  • Dropout7 は各 training step で activations のランダムな subset をゼロにするため、どの unit も特定の他の unit が存在することに依存できません。
  • Early stopping は validation loss を監視し、それが上向いたら止めます。
  • Data augmentation は、手元の training examples からさらに training examples を作ります。これは問題の根元を攻撃します。overfitting は parameters が多すぎることと同じくらい、データ不足でもあるからです。
  • Cross-validation はデータを kk 通りに分割し、kk 回学習します。held-out set を確保する余裕がないほどデータが少ないときに、test error の信頼できる推定を買う方法です。

ここで、その絵を壊す事実があります。

bias-variance の物語では、sweet spot を過ぎると parameters が増えるほど generalisation は悪くなる、と言います。現代の language models は、見るデータに対して古典的な規則が許すよりはるかに多くの parameters を持ち、それでも見事に generalise します。この2つの文はどちらも真であり、それらを両立させることがこの章で最も有用なことです。

40個の training points、20次元入力、ランダムな ReLU features、そして features の数 PP を2から5000まで sweep します。fit する解が多数あるときは、minimum-norm 解を選びます。

PPP/nP/ntrain RMSEtest RMSEw\lVert w \rVert
100.250.88221.25201.89
200.500.59621.16342.59
300.750.38961.53234.15
380.950.17693.716310.25
401.000.00005.814014.83
421.050.00003.16239.35
601.500.00001.10582.78
2005.000.00000.66380.98
150037.500.00000.58590.33
5000125.000.00000.56640.18

3つに分けて読みます。P/n=0.5P/n = 0.5 までは古典的な物語がそのまま成り立ちます。error は下がり、それから上がり始めます。P=n=40P = n = 40、つまりモデルがすべての training point を通るのにちょうど十分な parameters を持つ interpolation threshold では、test error が peak し、5.81 になります。小さいモデルの5倍悪い値です。この peak は古典的な警告であり、本物です。

その後、再び下がります。そして下がり続けます。P=5nP = 5n を超え、P=37nP = 37n を超え、P=125nP = 125n まで進みます。そこでの test error 0.5664 は、最良の under-parameterised model が達成したどの値よりも良いのです。40点に fit した5000 parameters のモデルが、この表で最良のモデルです。

これが double descent89 で、その仕組みは最後の列に見えています。いったん P>nP > n になると、training data に完全に fit する parameter settings は無限にあり、どれを得るかは選び方に依存します。minimum-norm 解は最小のものを選び、w\lVert w \rVert がその意味を示しています。threshold ちょうどでは14.83で peak します。そこでは interpolating solution が1つしかなく、どれほど極端でもそれを使うしかないからです。そして PP が大きくなるにつれて単調に下がります。parameters が多いほど、選べる interpolating solutions が増え、利用可能な最小のものが小さくなるからです。P=5000P = 5000 では norm は0.18で、threshold の80分の1です。

つまり、追加の parameters は complexity を増やしているのではありません。選択肢を増やしているのです。そして選択規則は、その選択肢を単純さに使います。正則化は損失関数の中にあるのではなく、algorithm の中にあります。小さな初期化からの gradient descent には small-norm solutions へ向かう documented bias があり、それがこの振る舞いが上の線形代数だけでなく、普通に学習された実際のネットワークにも現れる理由です。

実務上の帰結、そして Chapter 10 が依存するものはこれです。「モデルの parameters がデータより多いので overfit する」は有効な議論ではありません。モデルが threshold の左側にいた時代には良い規則でした。いま面白いものはすべて、そのはるか右側にあり、そこでは規則が反転します。

この章の道具があれば、表に入れられるデータ、つまり数値の行とラベルの列で機能するネットワークを学習できます。

言語はそれではありません。モデルが次の単語を予測できるようになる前に、そもそも「単語」とは何かを何かが決めなければなりません。その答えは文字でも単語でもなく、モデルが training data の raw bytes から学ぶ vocabulary です。学習が始まる前に一度だけ行われるその決定が、モデルが言えることの数、リクエストのコスト、そして法律試験に合格できるモデルが strawberry の文字数を確実には数えられない理由を決めます。

Chapter 7 では tokenizer を作ります。


上で使った residual connections については、He et al., Deep Residual Learning for Image RecognitionarXiv:1512.03385)を参照してください。Andrej Karpathy の Building makemore Part 3: Activations & Gradients, BatchNorm は、実際のモデルで activation-histogram diagnostic をたどるもので、この章の前半について最も手を動かして学べる扱いです。Yaser Abu-Mostafa の Learning From Data lectures 8 and 11–13 は、古典的な generalisation theory を、この章が1段落に圧縮した部分も含めて、きちんと説明しています。

  1. Glorot, X. and Bengio, Y. Understanding the difficulty of training deep feedforward neural networks. AISTATS (2010). 上の box で再現した variance-preservation の議論です。

  2. He, K., Zhang, X., Ren, S. and Sun, J. Delving Deep into Rectifiers: Surpassing Human-Level Performance on ImageNet Classification. arXiv:1502.01852 (2015).

  3. Ioffe, S. and Szegedy, C. Batch Normalization: Accelerating Deep Network Training by Reducing Internal Covariate Shift. arXiv:1502.03167 (2015). タイトルにある「internal covariate shift」という説明は、その後かなり大きく異論を受けている点に注意してください。層は機能しますが、なぜ機能するかについての元の説明は争われています。

  4. Ba, J. L., Kiros, J. R. and Hinton, G. E. Layer Normalization. arXiv:1607.06450 (2016).

  5. Kingma, D. P. and Ba, J. Adam: A Method for Stochastic Optimization. arXiv:1412.6980 (2014).

  6. Loshchilov, I. and Hutter, F. Decoupled Weight Decay Regularization. arXiv:1711.05101 (2017).

  7. Srivastava, N., Hinton, G., Krizhevsky, A., Sutskever, I. and Salakhutdinov, R. Dropout: A Simple Way to Prevent Neural Networks from Overfitting. JMLR 15, pp. 1929–1958 (2014).

  8. Belkin, M., Hsu, D., Ma, S. and Mandal, S. Reconciling modern machine-learning practice and the classical bias–variance trade-off. PNAS 116(32), pp. 15849–15854 (2019). この現象に名前を付けた論文です。

  9. Nakkiran, P., Kaplun, G., Bansal, Y., Yang, T., Barak, B. and Sutskever, I. Deep Double Descent: Where Bigger Models and More Data Hurt. arXiv:1912.02292 (2019). 実際の deep networks でこの効果を示し、model size axis だけでなく training time axis でも示しています。

モデル選びは、LIAにおまかせ。

すべてのAIモデルをひとつの場所で。今日から無料で。