損失と安定性 入門
分類器の最終層が出すのはスコアであり、そこから損失に至る各段は、算術がひそかに算術でなくなりうる場所だ。softmax、二乗誤差と交差エントロピー、 88.72 という指数の上限、それを消す最大値の減算、log-sum-exp、そしてどのフレームワークの損失も確率を受け取らない理由。このページの数値はすべて、隣の図が模擬 float32 の上で計算している —— だから上限の先まで押して、壊れるところを見ることができる。
スコアは確率ではない
分類器の最終層はクラスごとに1 つ数を吐く。その数のどこにも分布らしさはない。
分類器の最後の線形層は、クラスごとに実数を 1 つ返す —— logit だ。負でも正でも、巨大でも極小でもよく、何の制約もない。行列積の出力であり、行列積は確率という言葉を聞いたことがない。
ここに 5 つある。スライダーを引いて3 番目のスコアを動かし、この行が何をするか —— そして何を決してしないかを見てほしい:
下に出る合計が3 番目のスコアのドラッグにつれて漂うことに注目。静止時は 1.8、 にすれば 5.8 になる。どこかへ引き寄せる力は存在しないし、存在しようがない —— この数を作った層には正規化の手続きが入っていないからだ。
分布にするには、すべての値が正で、合計が 1 に固定されていなければならない。指数を取れば前半は手に入るが、代償がある ——3 番目のスコアを上げて 2 行目を見てほしい:
この行が形を失う速さを見てほしい。スコア 2 のとき最大の指数は 7.39、合計は 12.38。 では 408.42 のうち 403.43 で、残り 4 本は基線に潰れている。指数は数を正にするだけではない。スコアの差を比に変え、その比はスコア 1 単位ごとに e 倍になる。
後半は割り算 1 回で買える。各指数をその合計で割れば、元のスコアが何であれこの行は確率分布になる:
これで softmax は完成だ。指数を取り、和で割る。Σp はスライダーのどの位置でも 1.0000 と読める。分母が分子の数そのものの和だからだ。この不変条件は近似でも学習の結果でもない —— 算術であり、未訓練のネットでも成り立つ。
指数の前にスコアを定数で割ると、結果の鋭さを決めるつまみが手に入る。スコアを動かさないまま、温度を端から端まで引いてみる:
では勝者が小数第 4 位まで 1.0000 で、この図は argmax をしている。T = 4 では最大確率が 0.2905 で、5 つはほぼ横並びだ。スコアは一度も動いていない。温度はモデルの性質ではなく、サンプリング時に選ぶ「読み方」である。
つまり logit はスコアであり、確率はスコアの集合に対する読みだ。以降はすべてこの 2 文のあいだの隔たりの話になる。その渡り方が、訓練の出力を数にするか NaN にするかを決めるからだ。
2 つの損失、2 つの形
損失とは形だ。どの形を選ぶかが、モデルが大きく外したときの勾配を決める。
二乗誤差は誰もが最初に出会う損失で、回帰にはこれが正しい。ガウス雑音のもとでの最尤損失であり、勾配は 2(ŷ − y)、予測について凸だ。
お椀に沿ってマーカーを引き、損失を読み取ってほしい ——正解値は 1 の破線だ:
仕事をしているのは形そのものだ。罰は外れの二乗で増えるので、 2 外した予測は 1 外した予測の 4 倍のコストになり、勾配は誤差に比例して伸びる。遠くにいるほど強く押し戻される。
分類が問うのは別の問いだ。そこではモデルが正解クラスに確率を割り当て、罰したいのは「自信をもって間違う」ことだ。マーカーをゼロのほうへ引いてほしい:
左に立つ壁を見てほしい。モデルが確信して当てていれば −log p は 0、五分五分なら 0.69、p = 0.01 なら 4.61、そして p → 0 で発散する。これが one-hot 正解に対する交差エントロピー —— モデルが実際に出した答えの意外さであり、上界がないのは意図的だ。
同じ問いに二乗誤差を向けると、違いは微妙どころではなくなる。 2 本の曲線は同じ確率を読んでいる。マーカーを左へ、モデルが自信をもって間違っている領域へ引いてほしい:
で交差エントロピーは 3.00 を課し、二乗誤差は 0.90 しか課さない。しかも二乗誤差はどれだけ間違っても 1 で頭打ちなので、両者が最も食い違うのは、最も重要な場所そのものになる。
ただし訓練が使うのは損失の値ではなく勾配だ。そして二乗誤差はそこで完全に働かなくなる。マーカーを大きく外れたスコアまで引いてほしい:
二乗誤差は sigmoid の外側にあるので、勾配は σ(1 − σ) を抱え、モデルが最も誤る場所でちょうど死ぬ。、正解 1 で、交差エントロピーは 0.9975、二乗誤差は 0.00246 —— 405 倍弱い。飽和したユニットは学習をやめ、収束したように見える。
この打ち消しこそ、分類が交差エントロピーを使う理由だ。連鎖律の σ′ が、−log p の微分に現れる 1/p でちょうど消える。残るのは p − y だけで、その導出は第 06 節で見る。
指数は浮動小数点を使い切る
exp はここで最も速く伸びる関数で、落ちる先の形式は有限だ。
訓練中の数は IEEE-754 の浮動小数点に住む。binary32 は3.4028235e38 まで保持し、1.4012985e-45 より下は保持しない。この 2 定数が exp のぶつかる壁だ。
下の梯子は 1 段が 1 桁、最小の非正規化数から float32 の上限の先までを覆う。スライダーを引いてe のスコア乗が登るのを見てほしい:
1 段上がるごとに 10 倍なので、柱の高さは数の大きさではなく指数だ。 exp(2) は 1 のすぐ上に座る。 は上端を越え、上限より先ではすべての値が同じ 1 つの値、inf になる。緩やかな劣化はない —— exp(88.72) は有限の 3.39e38 で、exp(88.73) は無限大だ。
混合精度訓練は多くのテンソルを 16 ビットに移すが、float16 の指数部は 5 ビット、 float32 は 8 ビットある。スコアを上げて、どちらの上限が先に来るか見てほしい:
float16 の上限は 65,504、すなわち —— スコア 11 であって 89 ではない。行列積を半精度で回しながら、どの autocast 方針も softmax、log_softmax、cross_entropy を float32 に留めるのはこのためだ。掛け算は 16 ビットで安全だが、指数はそうではない。
では第 01 節のパイプラインをその壁に突っ込ませよう。最大のスコアを 88.72 の先へ上げ、3 行を同時に読む:
分子と分母が同時に溢れるので、勝ったクラスは inf ÷ inf を計算してNaN を得る。他のクラスは有限値 ÷ inf を計算し、きれいでもっともらしい 0 を得る。損失は NaN になる。逆伝播 1 回で、触れたパラメータもすべて NaN になる —— しかも何一つ例外を投げない。
これが本当に厄介な失敗様式だ。静かだからである。例外も警告も途中結果もない。実行は続き、損失はこの先ずっと nan を印字し、チェックポイントは無価値になる。あなたとこの結末のあいだに立つのが、次節の 2 行だ。
最大値を引く
恒等式 1 つで問題はまるごと消える。代償は行を 1 回余分に走査することだけだ。
任意のスカラー c について softmax(z) = softmax(z − c) が成り立つ。証明は 1 行だ。分子と分母に e−c を掛けても何も変わらず、 ez−c = eze−c だからである。これは信じるより見るほうがよい。
5 つのスコアを同じ量だけずらし、その下の確率の行を見てほしい:
下の行がまったく動かないことに注目。上の棒はすべて だけ動き、3 番目の確率は 0.5967 に留まる。 softmax が読むのはスコアどうしの差だけで、絶対的な高さは手に入れていない情報だ。
つまり c はこちらが選べる。そして特別な選び方が 1 つある。ずらす量を最大のスコアまで引き、指数を見てほしい:
c = 最大値のとき、ずらした後の最大スコアはちょうど 0 になり、その指数はちょうど 1、他はすべて (0, 1] に入る。 exp が渡される最大の引数がゼロなので、溢れようがない —— しかも分母は 1 以上なので、ゼロ除算も起こらない。
2 つの経路を 1 本の梯子に並べよう。生のスコアをスライダーの端まで上げ、どちらの柱が段の上に残るか見てほしい:
生の柱は 89 で梯子の上端を抜け、戻らない。ずらしたほうはどのスコアでも 1 と読める。ずらした後の最大の指数がつねに e0 だからだ。代価は最大値を探す走査 1 回、行を 3 回読むことだけだ。
だから softmax を定義そのままで書いてはいけない。exp(z) / exp(z).sum() は正しい式であり壊れたプログラムだ。最大値を引いた形は同じ式であり、失敗しようのないプログラムだ。どのフレームワークの softmax も後者である。
log-sum-exp と、決して起こらない log
ずらせば softmax は直る。だが log(softmax) は直らない。
交差エントロピーは −log p なので、訓練が欲しいのは対数確率だ。 softmax の後に log を取るのは 1 行であり、1 つのバグでもある。確率が浮動小数点の往復を生き延びねばならないからだ。
代わりに使う対象から。log Σ exp は柔らかい最大値 —— 最大のスコアに、差だけで決まる補正を足したものだ。補正を見てほしい:
差が 0 のとき補正は ln 2 = 0.6931 になる。2 つのスコアが等しいので和は最大値の 2 倍だ。差が になれば 0.0003 になる。 log-sum-exp は角を丸めた max であり、丸めは上位 2 つが接近している領域に閉じ込められている。
次に、その log が受け継ぐ失敗だ。2 つのスコアの差を広げ、2 位の確率が梯子を下るのを見てほしい:
1.1754944e-38 を下回ると float32 は正規化数の範囲を離れ、仮数のビットを削って粘る。 1.4012985e-45 を下回ると何も残らず、ちょうど 0を返す。それは で起きるが、 104 など何でもない —— 訓練済みの言語モデルは最上位と最下位の logit のあいだに 100 ナットを日常的に置く。
ちょうどゼロの確率の対数は −inf であり、素朴な経路が返すのはそれだ。下の 2 行は同じ量を計算している。差を 104 の先まで引いてほしい:
z − log Σ exp は確率を作らないので、表現する必要もない。 では −140.0 を返し、log(softmax) は −inf を返す。 −140 はまったく普通の float32 だ。引き算は正確で、損失を生む段は指数のほうだった —— それを飛ばしたのである。
log_softmax は log(softmax(x)) の包み紙ではない。数値範囲の異なる別の計算であり、それが独立したカーネルである理由だ。
生のスコアに指数を取ってはいけない
どのフレームワークでも損失が受け取るのは logits であり、確率を渡せば黙って学習を壊す。
logits からの交差エントロピーは logΣexp(z) − z[y] —— リダクションと引き算が 1 回ずつ、確率も割り算もない。この設計が元を取るのは勾配のところだ。
各 logit についての微分は p − y ——確率の行からone-hot 正解の行を引いたものだ。正解クラスのスコアを引き、上 2 行から 3 行目を読み取ってほしい:
いちばん下の行が、上 2 行の成分ごとの差そのものであることに注目。静止時、正解の列は 0.5967 − 1 = −0.4033 で、他の列はもとの確率のままだ。 σ′ は生き残らなかった。log から来る 1/p と指数から来る p がちょうど消えたからである。スコアが何であれ各成分は [−1, 1] に収まり、交差エントロピーの勾配がひとりでに発散しない理由もそこにある。
スコアの行ではなく確率の行を渡すと、損失はもう一度 softmaxをかける —— しかも黙って。確率もまた完全に正当な浮動小数点数だからだ。損失に渡すものを切り替えてほしい:
天井が現れるのを見てほしい。2 回目の softmax はすでに [0, 1] に押し込まれた入力を見るので、見られる最大の差は 1 であり、勝者は e / (e + n − 1) —— 5 クラスなら 0.4046 —— で頭打ちになる。モデルは訓練され、損失も下がる。ただしここでは 0.9048 より下には行けず、50,257 語彙なら 9.82 より下には行けない —— 一様に当てずっぽうを言った場合の 10.82 に対してである。
取引のもう半分は、融合形式には溢れるものが存在しないことだ。損失はずらす量に依存しえないので平らでなければならない。最大のスコアを 88.72 の先へ引き、どちらの経路がそれに同意するか見てほしい:
logΣexp が内部でずらしているので、融合経路は でも 0.5163 と読める。一方 −log(softmax) は 88.73 以降ずっと NaN だ。上限より下では小数第 4 位まで一致し、上では片方が数で片方は数ではない。
その上にメモリの話が乗る。8 × 1,024 位置、50,257 語彙の確率テンソルは float32 で 1.53 GiB、その勾配がさらに 1.53 GiB —— 毎ステップ確保し、書き込み、解放する。融合カーネルはどちらも作らない。 API が logits を受け取るのはそのためだ。
3 つの上限、5 行
式は 1 つ。それが数になるかどうかを決めるのは形式のほうだ。
このページ全体は、1 つの指数が 1 つの有限な形式に出会う話だった。スコアを引いて、どの形式なら結果を保持できるか見てほしい:
float16 は で力尽き、float32 は 88.72、 float64 は 709.78 —— この梯子の上端をまるごと越えている。 bfloat16 は float32 と同じ 8 ビットの指数部を持つので 88.72 まで届き、精度は 10 進 7 桁ではなく約 3 桁になる。
ここまでのすべてが 5 行に収まる。1 行目だけが教科書の定義に含まれない行であり、残る 4 行を安全にしているのもその 1 行だ:
m = z.max() # the shift
lse = m + log(exp(z - m).sum())
logp = z - lse # log_softmax
loss = lse - z[y] # cross-entropy
grad = softmax(z) - onehot(y) # p - y融合は数値の判断であると同時にメモリの判断でもあり、その代償は語彙とともに伸びる。小さな分類器から言語モデルまで引いて、 2 つの合計を読んでほしい:
では、融合したステップが抱える logits は 1.53 GiB、融合しないほうは 4.60 GiB。確率とその勾配が同じ形のコピーを増やすからだ。両軸とも対数なので 2 本は平行で、比はどの語彙でも 3 だ。
黙って失敗する 4 つ。生のスコアに指数を取ると、例外ではなくNaN が返る。softmax の出力に log を取ると、確率がアンダーフローした瞬間に −inf が返る。 logits を期待する損失に確率を渡すと、自信の足りないモデルが黙って訓練される。そしてどれかを float16 に移すと、コードを 1 行も変えないまま上限が 88.72 から 11.09 に落ちる。