微積分 入門
訓練ループが実際に回している分だけの微積分を、3 枚の絵の上に組み上げる —— 直線が触れている曲線、真上から見た碗の上の矢印、そして導関数が逆向きに辿っていく計算グラフ。感度、連鎖律、勾配、1 歩、そして 4 通りの壊れ方。すべての数値は隣の図が計算しているので、どこまで動かしても主張は崩れない。積分は出てこない。
傾きとは感度のこと
最適化器がいつも尋ねるただ 1 つの問いに答える数値 —— この入力をほんの少し動かしたら、出力はどれだけ動くのか。
訓練とは探索であり、その 1 ステップごとに、10 億個のつまみすべてに同じ問いが投げられる ——これを髪の毛 1 本ぶん回したら、損失は良くなるのか悪くなるのか、どれだけ動くのか。答えはつまみ 1 つにつき数値 1 つ。以下はその数値を安く手に入れるための仕掛けで、すべて 1 本の曲線の上に組み上げる。曲線の上で点をドラッグして、2 つの座標が追いかけるのを見てほしい:
関数とは 1 つの数値を別の数値に変える規則にすぎず、この点はその規則を 1 度見たものだ。まだ変化率ではない。変化率には読み取りが 2 回と、その間の距離が要る。h だけ先に 2 点目を取り、2 点を結べば、そこには計算できる傾きがある。刻み幅を縮めて、遠い端のすき間が閉じるのを見てほしい:
注目してほしいのは、読み値の落ち着き方の速さだ。h = 1.40 で割線は 2.05、 で 0.21、 で 0.05。割線は 0 に向かって歩いており、その極限が導関数 —— f′(1)、 2 つの量を名指ししたいときは df/dx と書く。同じ対象の 3 つの綴りで、論文にはどれも出てくる。
極限を取ると、割線は曲線にちょうど 1 点で触れ、そこで曲線と同じ向きを持つ直線になる。それが接線で、その傾きがその点の導関数だ。点をドラッグして、直線が一緒に傾くのを見てほしい:
どの x にも傾きが 1 つあることに注目してほしい。傾き自体が 1 つの関数になる:f′(x) = x² − 1。曲線が登るところでは正、下るところでは負、そして 2 本のグレーの目盛り —— と —— ではちょうどゼロで、曲線は平らになっている。最適化器が探しているのはこの平らな場所だ。
関数なのだから、絵に描ける。上が曲線、下がその傾き、入力軸は共有。どこかをドラッグすれば 2 つは一緒に動く:
よく見ると、下の曲線がゼロを横切る場所は、上の曲線が折り返す場所とぴったり一致する。ひとつ明確にしておくと、2 つの枠は同じ縮尺では描かれていない —— 傾きは角度ではなく読み値で読んでほしい。この絵が言っているのは、 1 本の式がすべての点の上昇速度を一度に運んでいるということ。割線を 1 本ずつ測るのではなく記号のまま導関数を求める価値は、そこにある。
接線については、もっと強い言い方があり、以降のすべてがそれに依存している —— 点に十分近づけば、曲線と直線は同じ対象になる。窓を半分にして、両者の最大のずれがどうなるか見てほしい:
見てのとおり、窓を半分にするたびに ずれはおおよそ 4 分の 1 になる —— で 0.112、 で 0.027 —— だからズームするほど線形モデルが勝つ。それが微分可能という言葉の中身のすべてだ:f(x + h) = f(x) + f′(x)·h + (誤差) で、誤差は h より速く死ぬ。あらゆるフレームワークのあらゆる勾配ステップは、小さな h の範囲でこの直線を信用しているだけのことだ。
だから壊れ方は、どれだけズームしても真っ直ぐにならない関数ということになる。 ReLU は原点でまさにそれで、しかもディープラーニングで最も使われる活性化関数だ。点をまで滑らせて、左右の片側傾きを読んでほしい:
折れ目では左からの答えが 0、右からの答えが 1 なので、接線は 1 本に定まらず、f′(0) は存在しない。しかも何もエラーを出さない:relu'(0) は PyTorch も TensorFlow も JAX も 0 を返す —— JAX は 2022 年に custom JVP で固定して以降で、それ以前は 1 だった。どのライブラリもこの食い違いを黙って決着させ、 1 つは途中で答えを変えている。[0, 1] のどの値も正当な劣勾配なので実害はないが、別の場所にバグを探しに行く前に知っておく価値はある。
つまみが一度にたくさん
モデルの入力は 1 つではなく数十億個ある。対処は、 1 変数の問いを変数の数だけ繰り返すこと —— そして「残りを止めておく」ことの代償に気をつけること。
パラメータが 2 つになると、損失は曲線をやめて地形になる —— どの組 (w₁, w₂) にも高さ L がある。真上から見るので、等高線が同じコストの組をつなぐ。点をドラッグして高さを読んでほしい:
これが二乗誤差損失の、最小点付近での本当の形だ。見るべき点は 2 つとも輪にある。輪は円ではなく楕円で —— 曲面は谷に沿う方向より横切る方向のほうが硬い —— そして急なところでは輪が密になる。中心では L = 0 に閉じていき、そこが訓練の探している最小点だ。
では w₂ を留め、w₁ だけを動かす。これは地形から曲線を 1 本切り出すことにあたり、その曲線の上では第 01 節の 1 変数の話に戻っている。断面に沿って滑らせ、それからどの断面にいるかを変えてほしい:
その接線の傾きが偏微分で、まっすぐな d ではなく丸まった ∂ を使って ∂L/∂w₁ と書く。この丸まりは新しい数学ではなく、どの変数を止めたかを読者に伝えるための注記にすぎない。計算も同じで、他の変数をすべて定数と見なし、第 01 節の規則を使うだけだ。
1 つの点に立つと、こうした問いは軸ごとに 1 つずつ、計 2 つあり、答えも 2 つある。点を動かして両方を読んでほしい —— 同じ点から出る 2 本の矢印がそれだ ——w₁ 方向の変化率とw₂ 方向の変化率:
最初の点ではひとつ目が −2.34、ふたつ目が 6.52:w₁ を上げると損失は下がり、w₂ を上げると急に上がる。D 変数の関数にはこれが D 個、変数 1 つにつき 1 つある。レシピは変わらず、増えるのは帳簿だけだ。
「止めておく」については引っかかりやすい点が 1 つあり、読むより手で触れてみる価値がある。∂L/∂w₁ 自体がすべての変数の関数なのだ。w₁ はそのままにして、もう一方だけを動かしてほしい:
見てほしい。w₁ には誰も触れていないのに、接線は傾いていく。w₂ = 1.60 では傾きは −0.76、 では 2.70、 では 6.16。w₂ = 1.25 では符号まで変わる。つまり勾配は点についての事実であって、パラメータについての事実ではない。ネットワークの他の重みが 1 つでも動けば、計算済みの偏微分はすべて古くなる。
変化率は掛け算になる
ネットワークとは関数が関数を食べていく構造だ。導関数を積み重ね全体に通す規則は、リンク 1 つにつき掛け算 1 回 —— 第 01 節より難しいものは最後まで出てこない。
リンクは 2 つ:x が g に入って u になり、u が h に入って y になる。下の 2 つの枠は中央の軸を共有しているので、左から入った揺さぶりは 2 回サイズを変えられて右から出てくる。揺さぶりを最初の幅から小さくして、 3 つの区間が一緒に縮んでいくのを見てほしい:
注目すべきは、どの数値が一致するかだ。最初の揺さぶりでは Δu/Δx = 1.80、Δy/Δu = 0.086 で、その積 0.155 がそのまま Δy/Δx の読み値になる。この等式は近似ではないし、揺さぶりが小さいことも要求しない ——Δu が約分されるだけで、端をつないだ 2 つの分数が共通項を約分するのと同じことだ。
揺さぶりをゼロまで縮めれば 3 つの比はそれぞれ導関数になり、等式は極限を無傷で生き延びる。下は 2 つの局所的な傾きを接線として描いたもの。入力をドラッグして、その 2 本が一緒に変わるのを見てほしい:
これが連鎖律だ:dy/dx = h′(g(x)) · g′(x)、約分が見える形で書けば dy/dx = dy/du · du/dx。最初の点では2 つの因子は 0.142 と 1.40、答えは 0.198。リンクが N 個なら因子も N 個で、レシピが難しくなることはない。
ただし罠が 1 つあり、誰もが一度は踏む:h′ は x ではなく順伝播が残した値で評価しなければならない。入力を まで動かすと図の読み値は 0.083。ところが h′(x)·g′(x) と書くと 0.190 になり、2.3 倍大きく、しかも黙って間違っている。正しい評価点を渡すのは順伝播だ。フレームワークが微分の前に順向きへ走る理由がこれである。
因子は掛け合わされるので、その大きさも複利で効いてくる。押しつぶす種類のリンクは 1 未満の因子を出す。ロジスティック曲線が典型例だ。曲線に沿ってドラッグして、下のその傾きを見てほしい:
下の曲線のピークはちょうど 0.25、場所は 。 ではもう 0.007 を下回る。この数値がすべてを語る:σ′ = σ(1−σ) は σ についての放物線で、頂点は σ = ½。だからどこのどの sigmoid 層も、 4 分の 1 より大きい因子を出すことはできない。
その因子を積み上げて、長い連鎖が変化率に何をするか見てほしい。各層が出すゲインと層の枚数を決めて、 1 つ 1 つのリンクが、そこに届いたものからどれだけ削り取るか見てほしい:
ゲイン 0.75 —— 軽い押しつぶし —— でも 24 層で千分の 1 になり、鎖が終わるずっと手前でビームは枠の底に貼りついてしまう。では 3.19×10⁻⁸。ゲインを sigmoid の最良ケース まで下げれば、 10 層だけで 9.54×10⁻⁷ しか残らない。これが勾配消失で、しかも黙って失敗する。エラーも NaN も出ず、ただ入力側の層の更新がゼロに丸められ、損失は動かない。逆にゲインを 1 より上に押せば同じ掛け算が爆発する —— こちらは NaN として自己申告してくれる。
すべての偏微分を 1 本の矢印に
変数ごとの答えをベクトルに詰めると、個々の成分にはなかった性質が生まれる —— 最も急な上りを指すのだ。
第 02 節の 2 つの偏微分を 1 本のベクトルにまとめると ∇L = (∂L/∂w₁, ∂L/∂w₂)、すなわち勾配だ。パラメータと同じ空間に住むので、地図の上に矢印として描ける。点を動かして、矢印と乗っている輪を見比べてほしい:
注目してほしい。矢印はどこでも輪と直角に交わる。これは偶然ではなく必然だ:輪に沿って動いても L はまったく変わらないので、その向きの変化率はゼロ、したがって勾配はその向きの成分を持たない。最初の点では ∇L は [−2.34, 6.52]、長さは 6.92 と読める。
これを最も急な向きと呼ぶのは主張であり、検証できる。任意の単位ベクトル û 方向の上昇速度は ∇L · û —— つまり勾配のその向きへの影だ。影をぐるりと一周させて、最大値を探してほしい:
影が最も長くなる場所に注目してほしい —— で、読み値は 6.92 —— 勾配自身の長さが、勾配自身の向きで出る。 では 0.03 —— 輪に沿う向きで、そこでは L はまったく変わらない。反対側の極で最も負になり、そこが最も急な下りだ。射影が射影される当のものの長さを超えることはない —— 証明はそれだけだ。
勾配がではないものが 1 つある。経路だ。無限小の 1 歩にとって最良の向きであって、最小点への方位ではない。碗が丸くなくなった瞬間に、2 つは袂を分かつ。碗を引き伸ばして、下り方向と最小点方向の角度を見てほしい:
では輪は円で、 2 本の矢印はぴったり重なり 0°。κ = 6 では 37° 離れ、 では 44°。最急降下は谷を下るのではなく谷を横切って歩き去る。そしてその数値 —— 曲面の最も硬い曲率と最も柔らかい曲率の比、すなわち条件数 —— が次の節の主題だ。
その 1 歩
たった 1 行 —— w ← w − η ∇L(w) —— と、 2 通りの間違え方。どちらもここで自分の手で起こせる。
勾配は上りを指し、こちらは下りたいので、引き算する。どれだけ動くかは微積分が決めることではない:導関数は極限でしか正しくなく、現実の 1 歩はどれも「接線がどこまで正直でいるか」への賭けだ。学習率 η がその賭け金にあたる。値を決めて、1 歩がどこに着地するか見てほしい:
見てのとおり、曲面は接線から離れて曲がっていくので、大きいほど良いとは限らない。L = 4.63 から出発して、η = 0.10 の 1 歩は 1.23 に着地する。 なら 0.53 で、ここからの 1 歩としてはこれが最良 —— 図に印がある。そこで着地点が並ぶ線は届く範囲でいちばん内側の等高線を突き抜けず、かすめるだけになる。値は η★ = ∇ᵀ∇ / ∇ᵀH∇ = 0.171。そこを過ぎると 1 歩はまた外へ登り、 では 0.64 に戻り、 では 5.52 —— 出発点より高い。
訓練はこれを何度も繰り返す。いまいる場所で勾配を計算し、1 歩進み、また繰り返す。実際に回して、トレードオフの両端を動かしてほしい ——刻み幅とステップ数:
注目してほしいのはジグザグだ。経路は谷を下るのではなく横切っていく。理由は第 04 節で述べたとおりで、そこへ追いやっているのは碗の形であって刻み幅ではない。η を小さくすると、振れ幅と同時に下降そのものも遅くなる。 12 ステップ後は、 で L = 0.223、η = 0.10 で 0.0606、 で 0.0036。どれもこの碗が課す上限より下なので、小さい刻み幅は安全なほうではなく、遅いほうにすぎない —— ここで最も速いのは で、 12 ステップで 6.6×10⁻⁴ に届く。
その上には硬い天井があり、これは好みの問題ではない。曲率方向ごとに 1 歩は誤差を 1 − ηλ 倍し、縮むのは |1 − ηλ| < 1 のときだけだ。最も硬い向きと最も柔らかい向きについてその倍率を描き、η をドラッグして 1 を横切る場所を探してほしい:
最初の学習率では、硬い向きは 1 ステップで誤差を ×0.40、柔らかい向きは ×0.90 にする。時間を食うのは柔らかい向きだ。硬い向きの曲線がちょうど 1 に達するのは 2/6 = 0.333。そこを越えると、その向きは 1 ステップごとに悪化する: なら 12 ステップで 9.92、 なら 1.31×10⁶。実際の訓練では数百ステップ後に損失が NaN になるあの現象で、大声で失敗してくれる診断の楽なほうの壊れ方だ。
静かなほうの失敗は、天井そのものが動くことだ。λ_max は曲面の性質なので、ある方向に細い碗はすべての方向に小さな η を強いる —— 本当は大きく踏み出したい方向にまで。碗を引き伸ばして、ステップ数を数えてほしい:
どの形でもその碗にとって最良の η で走らせている以上、ここで測っているのは条件数だけの効果だ —— 学習率はすでに調整済みという前提である。 なら 7 ステップ、κ = 6 なら 21、 なら 70 ——κ に比例して増える。実測された Hessian のスペクトルでは、訓練済みの画像分類器で最大固有値が数百のオーダー、残りはゼロ近傍(Ghorbani ら、2019)。現実の比は 6 どころではない。素の勾配降下を出荷する人がいないのはそのためで、モーメンタムも Adam も層正規化も、結局は碗を丸くする道具だ。
逆伝播
新しい数学は出てこない。第 03 節をグラフに適用し、安く済むほうの向きに走らせるだけだ。
中身のある最小のネットワークがこれだ:z₁ = w₁x、続いて a = tanh(z₁)、続いて z₂ = w₂a、最後に L = ½(z₂ − y)²。順に走らせて値を埋め、それから損失を起点に導関数をノード 1 つずつ遡らせてほしい:
逆向きの 1 ステップは、局所的な導関数を 1 回掛けることであり、その評価点は順伝播がそこに残していった値だ。∂L/∂z₂ = z₂ − y で 0.834、w₂ を掛けて ∂L/∂a = 1.334、tanh′ = 1 − a² を掛けて ∂L/∂z₁ = 0.407、x を掛けて ∂L/∂w₁ = 0.407。掛け算 4 回、第 01 節を超える微積分はゼロ。
コツは向きにある。連鎖律はどちら向きにも成り立つが、ネットワークには損失が1 つ、パラメータが数百万個ある。重み 6 個の同じグラフの上で、パラメータごとに 1 回と全部まとめて 1 回を比べ、スライダーを 1 目盛りずつ上げて、それぞれの代償を数えてほしい:
注目してほしいのは、6 回のドラッグが買ったものと 1 回が買ったものの差だ。順向きに進むと伝えているのは入力 1 つの影響なので、 1 目盛りごとに背骨の上へ経路が 1 本増え、答えは 1 つしか点かない —— n 個なら n 回。逆向きに進むと、伝えているのは出力 1 つの感度で、1 回の走査ですべての ∂L/∂wᵢ が一度に手に入る。70 億パラメータのモデルなら、逆伝播 1 回と順伝播 70 億回の差ということになる。
逆モードが代わりに払うのはメモリだ。順伝播の中間値はすべて、逆伝播がそこに到達するまで生かしておかねばならない。深さを決めて、全部保持する場合と少しだけ残して再計算する場合を比べてほしい:
なら素の逆伝播はテンソルを 48 個抱え、チェックポイント版は 14 個 ——√n 個の保存点と、再計算中の 1 区間ぶんの √n 個でピークは 2√n、計算量は約 30% 増(Chen ら、2016)。バッチサイズを決めているのはこの取引だ。
これで全部
6 つの規則、1 つのループ、4 通りの壊れ方。
d/dx [c] = 0 d/dx [eˣ] = eˣ d/dx [xⁿ] = n·xⁿ⁻¹ d/dx [ln x] = 1/x d/dx [f + g] = f′ + g′ d/dx [f(g)] = f′(g)·g′ loss = forward(w) # 中間値はすべて保持 g = backward(loss) # 1 回で全部の ∂L/∂wᵢ w -= eta * g # eta < 2 / 最大曲率
すべては 1 つの性質に乗っている —— 点の近くでは関数とその接線が 1 次より良い精度で一致する。それを壊すのは 4 つ ——折れ目、小さなゲインの積、2/λ_max を越えた刻み幅、条件数の悪い曲面 —— そして自分から知らせてくれるのは 3 番目だけだ。