自動微分 入門

loss.backward() はたった 1 行で、順伝播およそ 2 回ぶんの代金と引き換えに、 10 億個のパラメータすべての勾配を返す。このページはそれを分解する —— テープ、二つの掃引方向、メモリ、そして返る答えが意図した導関数でなくなる三つの場所。すべての数値は隣の図が計算している。

01

傾きを得る三つの道

モデルはつまみが 10 億個のプログラムであり、訓練はそのどれをどちらへ回せば損失が下がるかを知る必要がある。

書き下した関数から導関数を得る道は三つ。手で微分する、入力を動かして出力を見る、あるいはプログラムを微分する。三つ目が自動微分で、最初の二つがその存在理由である。

まず対象そのものを。導関数とは傾きだ。曲線に点を置くと、そこを通る接線が数を 1 つ持つ —— f がそこで登る速さである。点をドラッグして、その数が追ってくるのを見てほしい:

x = 2.00 · f′ = 28. 曲線に沿って点をドラッグする。矢印キーで 0.1 ずつ動き、Home で x = 2 に戻る
x = 2.00 · f′ = 28

注目してほしい。接線は、曲線に触れ、しかも沿って寝ている唯一の直線である。x = 2 で 28、f のほうは 49 —— 同じ 1 点についての別の事実で、訓練が欲しいのは傾きだ。

その数をどう得るかが問題である。定義が道を示す。右に h だけ離れた 2 点目を取り、2 点を通る直線を引き、h を縮める。間隔を詰めて、2 本が重なるのを見てほしい:

h = 1 · 32

見てのとおり、h = 1 では割線が 32、接線が 28。1 目盛ごとに差は縮み、ここでの誤差はちょうど 4h だ。なら h を機械の限界まで小さくすれば厳密になる?ならない。

以下の値はすべて本物の float64 であって、その模型ではない。h を より小さくすると、誤差は下がるのをやめ、また登る。両軸とも対数 —— 1 目盛が 10 倍である:

h = 1 · 0 桁
h = 1 · 0 桁

注目すべきは曲線が折り返す場所だ。左の腕が下がるのは割線がまだ弦だから。右の腕が登るのは、f(x+h) と f(x) の上位桁が一致していき、引き算がその桁を捨てるからだ。最良の刻み幅は 1e−8(√ε の近く)で、16 桁のうち 8 桁が正しくなる。 では x + h は x であり、答えはゼロになる。

桁が半分でもまだ生きられる。コストのほうは無理だ。 n 入力の最も素朴な関数 f = x₁·x₂·…·xₙ を取り、勾配全体に要る乗算回数を数える ——手で書く、有限差分で測る、逆向きに掃く:

n = 4 · 逆方向はまだ負けている
n = 4 · 逆方向はまだ負けている

よく見ると、手書きの勾配は成分ごとに積を組み直すのでコストは n(n−2)。有限差分は n+1 回の評価が要り、ほぼ同額で、しかも 8 桁目から間違う。逆向きの掃引は 3n−2 で、 n = 5 までは本当に負けている。 では 341 倍の差になり、言語モデルの n は 10 億の単位だ。

02

経路に沿った積

仕事はすべて 1 つの規則がやってのけ、その規則は誰もがもう習っている。

プログラムの微分は数式の微分より難しく聞こえるが、実際は易しい。プログラムはすでに、部品ごとの導関数を誰かが書き下してあるほど小さく分解されている。

これが y = (2x+3)² を 3 つの演算にしたものだ。値が各ワイヤを左から右へ流れ、どの箱の下にもその演算自身の導関数が、値の着いた場所で評価されて置かれている。x をドラッグしてほしい:

x = 2 · dy/dx = 28

注目してほしい。この 3 つの局所導関数が、このページ全体で唯一の微積分である。2 倍なら 2、定数の加算なら 1、2 乗なら 2v。どれも互いを知らない。連鎖律が言うのは、答えはその積 —— 2 · 1 · 14 = 28 —— であり、この積が x から y への経路だということだ。

積はどちらの順に掛けてもよく、以下のすべてがこの 1 つの事実から出る。左から数を運ぶ。1 から始め、演算を越えるたびにその局所導関数を掛ける。再生するか、1 演算ずつ歩いてほしい:

3 演算のうち 0 個を通過 · 1

見てのとおり、これが順方向モードで、運ばれる数は ẋ —— x が動いたときこのワイヤが動く速さである。値と同じ向きに並んで進むので、何も保存しない。

今度は同じ 3 回の掛け算を右から左へ。出力側で 1 から始めて後ろへ押すと、運ばれる数は x̄ —— このワイヤが動いたとき y が動く量 —— になる。再生するか、とよい:

3 演算のうち 0 個を通過 · 1

並べて読んでほしい。後ろ向きが 1、14、14、28、前向きが 1、2、2、28。同じ 3 因子、同じ答え、結合の順だけが逆だ。入力 1・出力 1 なら、2 つのモードはまったく互角である。

03

テープ

後ろへ戻るには、順方向のパスが何かを残していかねばならない。

順方向モードは振り返らない。逆方向の掃引は末端から始まるので、必要な局所導関数はすべて行きがけに計算済みであり、誰かが覚えていなければならない。

そこでフレームワークは、演算が走るたびに 1 行を書く。演算、入力、出力、そして逆方向規則が後で要るもの。プログラムを走らせてテープが伸びるのを見てほしい —— 最初は空だ:

テープは空

注目してほしいのは最後の列だ。+3 の逆方向規則は「勾配をそのまま通す」で、順方向から何も要らない。(·)² の規則は 2v なので入力を 1 つ残す。この列が自動微分のメモリコストのすべてである。

全体で Python 12 行 —— リスト、追記する演算、逆にたどるループ:

tape = []            # (out, ins, locals)

def mul(a, b):
    out = V(a.v * b.v)
    tape.append((out, (a, b), (b.v, a.v)))
    return out

def backward(y):
    y.g = 1.0
    for out, ins, loc in reversed(tape):
        for x, d in zip(ins, loc):
            x.g += out.g * d   # += , never =

逆方向はその for ループを逆順に回すだけだ。各行は届いた随伴変数に自分の局所導関数を掛け、結果を入力へ押し下げる。テープを下から読んでほしい:

シード ȳ = 1

不変条件であり、これが正しさの論証のすべてである。行は実行の逆順に再生されるので、ある値の消費者はその値が読まれる前にすべて終わっている —— つまり番が来たとき、その随伴変数にはすでに ∂y/∂v が v から y へのすべての経路について足し合わされて入っている。ループのアサーションとして:v.grad == sum(c.grad * dc_dv for c in consumers(v))。

その最後の列には代金があり、演算ごとに値段が違う。 8 × 1024 × 768 の fp16 活性 1 枚 —— GPT-2 の形、12.00 MiB —— を単位に測ると、規則が残さねばならないものは 1 桁以上の幅を持つ。 4 つを切り替えてほしい:

add · 0 B

見てのとおり add はタダで、規則が定数だからである。relu は入力の符号だけでよい —— 要素あたり 1 ビット、0.75 MiB。square は入力を残して 12.00 MiB、 は両方を残して 24.00 MiB。乗算を加算に替えるのは、この表の取引でもある。

まっすぐな鎖が隠していたことが 1 つある。x を 2 回使う —— y = x²·(x+1) —— と、x から y への経路が 2 本になる。x をドラッグして、値が出ていき、各経路が持ち帰るものが戻り、合流するノードがそれをどう扱うかを見てほしい:

x = 2 · x̄ = 16

よく見ると、ノードは足す。x = 2 では 2 本が12 と 4 を持ち帰り、x̄ は 16 —— 3x² + 2x —— になる。この += があるから、誰も経路を数え上げないのに経路の総和が正しく出る —— そして §06 では、これが 1 行の書き忘れで訓練を壊す理由になる。

04

1 回の掃引、1 つの問い

掃引は安い。ただし万能ではない。

ここまではすべて入力 1・出力 1 で、2 つのモードが引き分ける唯一の形だった。入力 2・出力 3 のプログラムに替えると、その導関数は数ではなく 3 × 2 のブロックになる。出力と入力の組ごとに 1 マスだ。

順方向モードが一度に運べるシードは 1 本。 ẋ₁ = 1、ẋ₂ = 0 と置けば、1 回の掃引は x₁ が動いたときの 3 出力の動きを教え —— x₂ については何も教えない。シードを切り替えて、どのワイヤが灯るかを見てほしい:

ẋ₁ = 1 · 1 列ぶん

見てのとおり、シード 1 本が生むのは 3 つの数、つまりブロックの 1 列である。この 3 × 2 のヤコビ行列には順方向が 2 回、入力 1 つにつき 1 回要り、両方を一度に取る巧いシードは存在しない。掃引はシードについて線形で、独立な 2 列には独立な 2 本のシードが要る。

大きくすれば、それがコストモデルのすべてだ。これは 4 × 6 のヤコビ行列 —— 入力 6、出力 4 —— で、順方向 1 回につき 1 列が埋まる。スライダーを押して埋めてほしい:

6 回のうち 0 回

注目してほしい。6 列に6 回で、回数は入力の個数であり、出力の個数はこの式に入らない。出力 4 のプログラムと 400 のプログラムは、順方向モードには同じ値段である。

§02 の逆方向の掃引は、同じ絵を 90 度回したものだ。出力を 1 つ選んで 1 を置き、後ろへ押すと、その出力がすべての入力にどう応じるかが一度に分かる —— 取れるのは列ではなく行である:

4 回のうち 0 回

今度は 6 回ではなく4 回。回数がいまや出力の個数だからだ。最適化は何もしていない。同じ連鎖律、同じ局所導関数、掃引あたり同じ乗算回数。変わったのは結合の順と、それに伴ってどちらの次元に払うかだけである。

05

逆方向の掃引が勝つ理由

訓練とは、数が 10 億入って 1 つ出てくることだ。

損失はスカラーなので、loss.backward() が求めているのは 1 × n のヤコビ行列 —— 1 行、n 列、n はパラメータ数 —— である。それを埋める 2 通りの値段は §04 で出してある。

下の 2 つのブロックがそのヤコビ行列で、左は 1 列ずつ、右は 1 行ずつ埋まる。入力と出力の個数を決めて、どちらが先に埋まるかを読んでほしい:

n = 8 · m = 1

見てのとおり、入力 8・出力 1 なら逆方向のブロックは 1 回で終わり、順方向のブロックは 8 回要る。形をひっくり返して に出力を多くすれば、同じ論法で順方向が勝つ。ブロックの短い辺に沿って掃け。

損失についてはどちらが短いかに疑いはない。出力を 1 に固定して、各モードに要る掃引回数を入力の個数に対して描く。両軸とも対数なので、まっすぐな対角線はそのまま比例を意味する:

n = 1
n = 1

逆方向は 1 に貼りついた水平線、順方向が対角線である。では 1 回対 100 万回。 70 億パラメータのモデルなら、逆方向 1 回が与えるものに順方向は 70 億回要る。

掃引はタダではないが上から抑えられている。Griewank の「安い勾配」は、入力が何個でも逆方向の勾配を関数評価 4 回以内に抑える。 Transformer での実測は順方向 1 に対して逆方向 2 程度だ。パラメータ数を動かして、棒が動かないのを見てほしい:

125M · 3×

注目してほしい。訓練 1 ステップは順方向の 3 倍で、 4 倍を超えないことが証明されている。1.25 億でも 4050 億でも同じだ。訓練 1 トークンあたり 6·N·P FLOPs、推論で 2·N·P はここから出る —— 偶然ではなく算術である。微分する対象が増えても勾配は高くならない。

高くなるのは別のところだ。逆方向の最初の 1 行が読まれる前に、テープは順方向パス全体の保存済み活性を抱えていなければならず、メモリは深さとともに増える —— 一部を捨てて再計算しない限りは。96 層のスタックで、k 層ごとの境界だけを残す:

k = 1
k = 1 · 97 / 96

よく見ると、チェックポイントなしのピークは 96 層ぶん。 でチェックポイントすれば 20 —— 4.8 分の 1 —— になる。 ⌈96/k⌉ 個の境界を保存し、再生中だけ k 層を生かすからで、和は √96 の近くで最小だ。代金は順方向 1 回ぶん、1 ステップあたり 3 割の計算量。「逆方向モードは 1 回」の二段目の答えは、時間が 1 回、空間がテープ 1 巻、である。

06

厳密、ただし何について

自動微分は最後の 1 ビットまで厳密である。だからこそ、何について厳密なのかを正確にしておく値打ちがある。

ここまでの数値には打ち切りも桁落ちもなかった。しかしテープが微分するのは実際に走ったプログラムで、それと書きたかった関数は、いつも同じ対象とは限らない。

最もよく使う活性化関数から。relu は 0 で接線を持たない —— 片側の傾きは 0、もう片側は 1、折れ目には定まる直線がない。折れ目まで点を引きずり、フレームワークが返すものを読んでほしい:

x = 1.40 · relu′ = 1. 曲線に沿って点をドラッグする。矢印キーで 0.1 ずつ動き、Home で x = 1.4 に戻る
x = 1.40 · relu′ = 1

見てのとおり、約束ごととして 0 を返し、警告は出さない。0 も 1 も ½ も擁護できる劣勾配で、 PyTorch は 0 を選ぶ。ほぼ絶対に問題にならない —— 浮動小数がちょうど 0 に乗る確率は無視できる —— が、誰かが x - x から規則を書き、勾配がどこへ消えたのかと悩む日には、途方もなく問題になる。

深刻なほうは制御フローだ。Python の if は実際に走った枝だけをテープに載せ、ほかは載せない。だから判定そのものは行にすらならず、導関数を持たない。点を段差の向こうへ引きずってほしい:

x = 1.40 · y = 1 · dy/dx = 0. 点を段差の向こうへドラッグする。矢印キーで 0.1 ずつ動き、Home で x = 1.4 に戻る
x = 1.40 · y = 1 · dy/dx = 0

表示が動かないのを見てほしい。点が越えるとき y は 0 から 1 になるのに、両側とも dy/dx は 0 と出る。比較が一度も記録されなかったからだ。これは静かに失敗する。モデルは訓練され、損失は別の理由で下がり、離散的な判断は計算したどの勾配からも見えていなかった。

対処は、わざと別の関数を微分することだ。段差をσ(x/τ) に置き換える。どこでも勾配を持ち、代わりに温度を 1 つ払う。τ を下げて勾配を見てほしい:

τ = 1

よく見ると、τ が下がるほど曲線は代役の段差に近づき ——その勾配は幅およそ τ の尖った峰に潰れる。 ではほとんどのサンプルがその外に落ちて何も学ばないほど細い。ストレートスルー推定器も Gumbel-softmax もこの取引を管理する方法で、取引は消えない。

最後の 1 つは落とし穴で、それは §03 の += そのものである。勾配バッファは設計上たまるので、意図に関係なくたまる。まずその行なしで 4 バッチ走らせ、次にオンにしてほしい:

バッチ 0 まで · 0×

注目してほしい。どちらのループもエラーなく走る。zero_grad() なしでは .grad にそれまでの全バッチの和が入るので、バッチ 4 は 4 倍大きすぎる 1 歩を踏む —— 落ちはせず、ただ訓練が下手になる。いちばん高くつく種類のバグだ。このページの 3 つの失敗のうち 2 つはこう静かで、 3 つ目の保存済みテンソルの上書きは大声で失敗する。テープのバージョンカウンタが捕まえるからだ。

07

通しで走らせる

技術の両半分を、端から端まで、1 つの操作子の下で。

§03 の 12 行を、§03 のひし形の上で走らせる。順方向に 3 演算がそれぞれ 1 行を追記し、続いて逆方向に 4 回の辺の再生がそれを読み取る。再生するか、1 ステップずつ歩いてほしい:

まだ何も走っていない

見てのとおり、3 行のテープから dy/dx = 16。演算ごとの規則は locals の 2 つの数だけで、このページのほかのどこにも微積分はなく、刻み幅も桁落ちもない。

本物のフレームワークが足すものはすべて、まさにこの上の工学である。規則を Python でなく C++ で書き、浮動小数でなくテンソルを扱い、分岐したらリストでなくグラフを使い、テープがメモリに収まるよう §05 のチェックポイントを入れる。