誤差逆伝播 入門

誤差逆伝播はネットワークを 1 回逆向きに辿るだけで、すべてのパラメータの勾配を手に入れる。そしてそれは 1 枚の絵からできている —— 曲線と、その上を滑る線。その線の角度が勾配だ。主張は 3 つ、どれも隣の図が計算している —— 逆伝播は傾きの積であること、その積のおかげで 100 万回ではなく 1 回で済むこと、そして同じ積のせいで深層ネットワークが消失し爆発すること。

01

勾配とは傾きのこと

比喩ではない。線の実際の急さであり、手でつかんで回せる。

この primer が使うネットワークはこれ 1 つ:入力 x = 2、 ReLU の隠れニューロン 1 個、出力 1 個、目標 y = 5 に対する二乗誤差、そして w₁ = 1.5、b₁ = 0.5、w₂ = 0.8、b₂ = 0.1 から始まる 4 つのパラメータ。予測は 2.90、損失は 4.41。どれか 1 つを動かせば、この数は曲線を描く —— 1 つ選んでほしい:

w₁ に沿って、損失 4.41

4 本の曲線が、どれも同じ 1 点を通る。w₁ の曲線には −0.25 に角がある。そこで隠れユニットの ReLU が切れ、損失は w₁ にまったく反応しなくなる。b₂ の方はきれいな放物線だ。逆伝播が返すのはパラメータ 1 つにつき 1 個の数で、そのどれもがそのパラメータ自身の曲線の性質だ。

その性質とは、いま立っている点の急さのこと。w₂ の曲線を取り出し、その上に接線を置いて、点をドラッグしながら傾きが回るのを見てほしい:

w₂ = 0.80、勾配 −14.70
w₂ = 0.80、勾配 −14.70

接線がそのまま勾配であり、ほかに計算するものはない。初期の重みでは −14.70 の下り傾き、これがちょうど ∂L/∂w₂ だ。右へ滑らせて まで行くと平らになる。損失が最小だから傾きは 0。「勾配がない」とはこの見た目のことだ。

では −14.70 はどこから来るのか。微分の定義からだ:近い 2 点での損失を、その距離で割る。間隔を縮めて、割線が接線に重なっていくのを見てほしい:

間隔 1.60、割線の傾き 4.90

初期の間隔 1.60 での符号に注目してほしい。割線は +4.90、真の値と逆向きだ —— それだけ広い一歩は最小値を跨いでしまう。 まで縮めると −14.09 になる。torch.autograd.gradcheck がやっているのはこれだ。 1 パラメータごとに順伝播 1 回ぶんかかるので、テストにはなっても訓練ループにはならない。

傾きの役目は、どちらへ動くかを言うこと。上り方向を指すので、逆向きに −η · ∂L/∂w₂ だけ進む。学習率は 0 から始まるので歩幅はまだゼロ。上げていくと、重みが曲線を滑り降りる:

η = 0.000、1 歩あとの損失 4.41

を越えたところを見てほしい。一歩が最小値を越えて、向こう側の壁を登り始める。 ではちょうど出発点と同じ損失に着地し、それより先はどの一歩も状況を悪くする。この閾値は 2/L″、この放物線なら 1/a² = 0.082 —— 勾配が教えるのは向きであって、距離ではない。

傾きを鎖につなぐ前に、もう 1 つ。逆伝播が必要とするのは損失の傾きだけでなく、すべてのノードの傾きだ。そして線形ノードならそれはただでついてくる。出力の重みを動かして、このノードの直線が入力を軸に回るのを見てほしい:

w₂ = 0.80、局所の傾き 0.80

ŷ = w₂·a + b₂ は直線なので、その傾きはどの入力でも同じで、値は w₂ —— 順伝播がすでに手にしていた数だ。どの演算の局所微分もこの形をしている:順伝播がどうせ計算した値だけでできた短い式。逆伝播が安いのは、この式が安いからだ。

02

連鎖律とは傾きの掛け算

変化率は掛け算される。その 1 つの事実をグラフのすべての辺に当てはめたものが、アルゴリズムのすべてだ。

ネットワークは小さな関数の積み重ねで、それぞれが次に渡していく。入力を少し動かすと中間が 2 倍動き、中間を少し動かすと損失が 0.8 倍動くなら、入力を少し動かしたとき損失は 0.8 × 2 = 1.6 動く。それ以上の仕掛けはない。

この絵はその掛け算を目に見える形にしている。同じゆらぎが 3 回描かれる —— 入力の位置で 1 回、1 段目の傾きのあとに 1 回、2 段目のあとに 1 回 —— 各レールの長さは、そのすぐ上のレールにその段の傾きを掛けた長さになる:

1 段目の傾き 2.00 · 2 段目の傾き 0.80

どちらかのスライダーが を通るとどうなるかに注目してほしい。もう一方がいくら大きくても、いちばん下のレールは丸ごと消える。1 段が死ねば経路全体が死ぬ —— これが §05 の問題のミニチュアだ。両方を 1 より大きくすれば、下のレールは上より長くなる。勾配は戻る途中で縮むのと同じくらい簡単に伸びる。

このネットワークはそういう段が 4 つ:x → z → a → ŷ → L。何かが戻り始める前に、まず順伝播が走ってすべての線に値を置いておく必要がある。再生を押すか、トラックをドラッグして、4 つの値が順に現れるのを見てほしい:

ステップ 0 / 5 —— x = 2.00

これらの数はどれももう一度使われる。活性値 a = 3.50 は ŷ の w₂ についての局所微分になり、x = 2 は z の w₁ についての局所微分になる。順伝播が中間結果を捨てないのはそのためで、§03 はその判断に値札をつける。

では同じグラフを逆向きに。逆伝播の各ステップは、直前の結果に局所微分を 1 つ掛けて左へ渡すだけだ。トラックをドラッグして、勾配が 1 本ずつ線の下に現れるのを読んでほしい:

ステップ 0 / 5 —— L での勾配 = 1.00

矢印の上の乗数を見てほしい。−4.20、次に × 0.80、次に × 1、次に × 1.50。a での勾配は −3.36 で、 ReLU を通ったあとも −3.36 のまま。z = 3.5 は正で、そこでの ReLU の傾きはちょうど 1 だからだ。逆伝播は何かを導出し直すことはない。表を辿って掛けるだけだ。

安く済む理由は使い回しにある。どのパラメータもいずれかの線にぶら下がっていて、その線の勾配は全員ぶんまとめて 1 回だけ計算される。 4 つのパラメータを順に送って、経路の共有部分が動かないまま掛け算が 1 回だけ増えるのを見てほしい:

b₂:共有部分の上に掛け算がさらに 1 回

別々に計算すれば、この 4 つの勾配には 2 + 2 + 4 + 4 = 12 回の掛け算がいる。まとめれば 7 回:線の勾配 3 つと、パラメータごとに 1 回ずつ。実際のモデルの規模ではこの比がすべてだ —— 96 層のは単独なら 96 層を歩き直すが、実際には 1 つ上の層が埋めた線の上で掛け算 1 回ぶんしか払わない。

03

逆向きに進むのは、その方が安いから

連鎖律はどちらの端から始めろとは言わない。決めるのは問題の形だ。

掛け算の列は左からでも右からでも評価できる。入力側から始めればある入力の影響を前へ運び、損失から始めれば損失の感度を後ろへ運ぶ。どちらも正しい。安いのは片方だけで、どちらかは形が決める。

こちらの形は、入力が数百万個のパラメータ、出力がスカラー 1 個。パラメータ数を増やして、順方向モードの走査が積み上がる一方、逆方向モードの走査が 1 回のままなのを見てほしい:

順方向 · 走査 4 回

順方向モードは入力ごとに 1 回、逆方向モードは出力ごとに 1 回。なら 12 回対 1 回で、1750 億パラメータのモデルなら 1750 億回になる。訓練が成立する理由はここに尽きる。損失はただの 1 個の数なので、安い側の端は遠い方の端だ。

とはいえ逆方向モードもただではない。順伝播が置いていったあとの線をすべて見に行くので、順伝播は中間結果を捨てられない。深さを下げてまた上げ、GPT-3 の形で溜まるテンソルの山を見てほしい:

96 層、2,048 トークン

GPT-3 の 96 層は、2,048 トークンの系列 1 本で275 GB の活性値を抱える —— カード 1 枚が 80 GB の世界で、だ。この数字は Korthikanti らのもので、彼らの式はその内訳も示す。82.1 GBが系列長に比例する部分、193 GBが系列長の 2 乗に比例するアテンション行列だ。

効いてくるのはその 2 乗の項だ。下の図は両軸とも対数なので、直線はべき乗則を意味し、直線の急さがその指数になる。コンテキスト長を送っていき、合計が線形の項から離れていくのを見てほしい:

2,048 トークン、うちアテンションが 193 GB

では 2 本の差は 2 倍以内だが、では合計 13,027 GB に対し線形項は657 GB。長いコンテキストはまず記憶の問題だ —— FlashAttention はそのために存在し、この項が測るアテンション行列をそもそも実体化しない。

もう 1 つの逃げ道は、ほとんど何も残さず、作り直す代価を払うこと。各層の入力だけをチェックポイントにしておき、逆伝播は活性値が必要になる直前にその区間の順伝播を走らせ直す。保存と再計算を切り替えてみてほしい:

96 層、保存

逆伝播の演算量は順伝播のおよそ 2 倍なので、順伝播をもう 1 回走らせても仕事は 3 分の 1 ほど増えるだけ —— その代わりに275 GB が7.7 GB まで落ちる。だから勾配チェックポイントはどの訓練フレームワークでも 1 行のフラグで、モデルが載らないときに最初に入れるものになっている。

04

局所微分はどこから来るのか

あの逆向きの道のりに出てくる乗数はどれも、どこかの曲線のどこかの点の傾きだ。その曲線がこれ。

線形層が差し出すのは自分の重みで、入力によらない。活性化関数は違う。その傾きは順伝播がたまたまどこに着地したかで決まるので、同じネットワークでも例が違えば返ってくる乗数は違う。

活性化関数 4 つ、横軸は共通。関数を選んで、接線を曲線に沿ってドラッグしてほしい。読み値が、逆伝播がこれから掛ける数だ:

σ、z = 0.00、傾き 0.250
σ、z = 0.00、傾き 0.250

シグモイドが真ん中以外どれだけ平らかに注目してほしい。いちばん急なのは原点で、そこですら傾きは 0.250 しかない。tanh の頂点は 1.000、ReLU は生きている側でちょうど 1、反対側でちょうど 0、 GELU は 1 を少し超えて 1.129 で頭打ちになる。 100 層のネットワークの運命は、この 4 つの数で決まる。

シグモイドの天井は単独で描く価値がある。曲線の下にあるのがその傾きで、同じ活性化前の軸に対して描き、天井を横線で引いてある。カーソルを動かして、同じ数を 2 か所で読んでほしい:

z = 0.00、傾き 0.250

この傾きの曲線は σ(z)·(1 − σ(z))、つまり足して 1 になる [0, 1] の 2 数の積だ。だから 2 つが等しいとき、すなわち で最大になり、その値はちょうど 4 分の 1。軸上のどの入力を持ってきても、シグモイドが 0.25 より大きい数を返すことはない。これは傾向ではなく上界で、 §05 はそれを自分自身と掛け合わせる。

同じユニットを裾に押し込むと、その上界はもう関係なくなる。実際の値がはるか下にあるからだ。点を外へ滑らせて、傾きの三角形がぺしゃんこになるのを見てほしい:

z = 0.00、傾き 0.2500

では傾きは 0.0025 —— 最大値の 100 分の 1 —— で、 では 0.0003。このユニットは壊れておらず、自信たっぷりに 0.9997 を出し続ける。ただ下流のすべてが 1 万分の 3 を掛けられるので、学習できなくなっただけだ。これが飽和、静かな失敗で、損失がただ動かなくなる。

ReLU には飽和する裾がないが、もっとたちの悪い手口がある。死んでいる側の傾きはちょうど 0 なので、バッチ全体が折れ点の左に落ちたユニットは何も受け取らない。バイアスを下へドラッグして、バッチが渡っていくのを見てほしい:

バイアス 0.00、8 個中 6 個が生きている
バイアス 0.00、8 個中 6 個が生きている

8 個の入力のうち、はじめは 6 個が生きている。 を越えると 1 個も残らない。どの例も傾きは 0、したがって ∂L/∂w はちょうど 0、重みは動かず、バイアスも動かず、このユニットは永久に死ぬ。例外は何も飛ばない。バッチをそこに留めたまま、ユニットが計算する中身を変えてみてほしい:

ReLU、バイアス −2.50

ReLU では 8 個の傾きの和はちょうど 0.000。leaky ReLU の負側の定数 0.01 はそれを 0.080 にし、GELU では 0.152 になる。どれも大きくはないが、どれもゼロではない。重要なのはその一点で、何かを受け取れるユニットはまだ這い上がれる。

05

50 個の傾きを掛け合わせる

逆伝播の不変条件は、第 i 層の勾配が第 i+1 層の勾配に局所微分を 1 つ掛けたものであること。繰り返せば、手元に残るのは積だ。

深さの難しさはこれで言い尽くされている。50 層とは50 個の係数の積であり、多数の数の積は和のようには振る舞わない。ゆっくりずれるのではなく、複利で効く。2 つの破綻の仕方は、積にできる 2 つのことそのものだ。

下の棒は各層での勾配の大きさ、軸は対数 —— 1 目盛が 10 倍なので、まっすぐ並んだ棒は指数的な増減だ。どの段の係数もちょうど 1.00 から始まるので縮みも伸びもしない。係数をどちらかへ動かしてほしい:

係数 1.00、50 層を経て 1.0e+0

必要なずれがどれほど小さいかに注目してほしい。 —— 1 層あたり 1 割の目減り、無害に聞こえる —— でも 50 層で 5.2e−3 まで落ちる。 なら第 50 層に届く勾配は 1.3e−20、 なら 1.6e+10。積をそのままにする係数はちょうど 1 つしかなく、偶然そこに座るネットワークはない。

シグモイドの場合、係数は選べない。 §04 で見たとおりその傾きは 4 分の 1 を超えられないので、シグモイドを積むと「どれだけの勾配が生き残れるか」に保証された上限がつく。層を増やして、半精度が音を上げる 2 本の線を棒が通り抜けるのを見てほしい:

4 層、3.9e−3 に到達

この 2 本はどちらも 2 のちょうどのべき乗なので、整数の層数に着地する。0.25⁷ = 2⁻¹⁴ は fp16 の最小の正規数、つまりが半精度の勾配がビットを失い始める地点だ。0.25¹² = 2⁻²⁴ は fp16 に残された最後の数なので、で形式を使い切り、 13 層でちょうど 0 になる。しかもこれは、どのユニットも最も急な点にいる最良のケースだ。

これにはちょうど 1 回の掛け算で済む定石がある。逆伝播を呼ぶ前に損失を大きくしておけば、グラフのすべての勾配が同じ倍率で返ってくる。この列を床から持ち上げてほしい:

ロススケール 1

ロススケール では、 12 層ぶんの走りは 9.8e−4 に着地する。 fp16 の正規数の範囲に収まり、更新の前に 16,384 で割り戻すので、歩幅は変わらず、変わったのは表現だけだ。2 本目の線にも注意してほしい。いちばん上の棒はfp16 自身の最大値より下にいる必要がある。GradScaler が倍々に上げてはあふれたら戻すのはそのためだ。

上向きの壁は、思われているより硬く、そして近い。ここでは鎖は 128 層で、どの段も 1 より大きい数を掛ける。上に横線で引いてあるのがfloat32 の最大値だ。棒がそこにぶつかるまで係数を上げてほしい:

係数 1.60、1.3e+26 に到達

—— 1 層ごとに倍 —— では、積は第 128 層で 3.4e+38 を越え、それより先の棒はすべて inf と読める。勾配の inf は重み更新の inf になり、inf − inf は nan。2 反復あとには全パラメータが nan で、損失は数字を印字しなくなる。爆発は派手だ。それが唯一の救いでもある。

直す前に 1 つ正直な訂正を。層の係数はスカラーではない —— ヤコビ行列であり、方向ごとに伸ばし方が違う。勾配が到来する向きを回して、出ていく側が楕円の上を走るのを見てほしい:

120°
σ_max = 1.00

円は勾配が到来しうるすべての向き、楕円はこの層がそれらを送る先だ。長半径が σ_max、短い方が σ_min で、ここではその 0.63 倍。つまり「係数 1.00」とは、勾配が向き次第で 0.63〜1.00 を掛けられるということだ。 1 の近くにいる必要があるのは σ_max の方で、スペクトル正規化が縛るのはそれだ。

06

重みはどこから始まるか

逆伝播は渡されたものを改善する。初期の重みの渡し方が 2 通り悪いと、何ひとつ改善しなくなる。

最初に思いつくのは、すべての重みを 0にすることだ。対称で、偏りがなく、1 行で書ける。そして最初の 1 歩の前にネットワークを殺す —— 損失ではなく、逆伝播そのものの性質として。

隠れニューロン 4 個を、それぞれの入力重みベクトルとして描き、先端に受け取る更新をぶら下げてある。最初は 4 本とも完全に重なっている。広がりを開いて、 1 本だったところに 4 本が現れるのを見てほしい:

広がり 0.00、区別できるニューロンは 1 個

同一の重みは同一の入力を見るので同一の出力を出し、同一の勾配を受け取り、更新後もやはり同一のままだ。逆伝播は渡された対称性を保つだけで、それを破る仕組みを持たない。こう初期化された 100 ニューロンの層は1 個のニューロンが 100 個、永久にそうだ —— しかもエラーは出さず、 1 ニューロンぶんの精度で頭打ちになるだけだ。

だから重みはランダムに始める。ただしそのランダムさのスケールもただではない。§05 の積は、逆向きの勾配と同じだけ、順方向の活性値にも効くからだ。下は 12 層の ReLU スタックの層ごとの分散で、初期化は He のレシピにゲインを掛けてある。ゲインを動かしてほしい:

ゲイン 1.00、12 層目の分散は 1.0e+0

ゲインがちょうど 1.00 のとき分散は平らだ。どの層も受け取ったものをそのまま渡す。 なら 12 層目で 1.9e−4、 なら 3.2e+3。ハイパーパラメータ 1 つの 3 割のずれが、12 回複利で効く。初期化に名前つきのレシピがあってデフォルト 1 つで済まないのは、これが理由だ。

レシピどうしの差は 2 倍ぶんで、その 2 は ReLU が捨てる半分そのものだ。下は、ある層に届く 64 個の活性化前の値。左側がReLU が 0 にする半分だ。層番号を送ってから、レシピを切り替えてほしい:

1 層目、Xavier

ReLU は分布の半分を消すので、層から出る分散は入る分散の半分になる。だからXavier の √(1/n)は毎層で信号を半分にし、の広がりは 0.022、出発時の 1.000 に対してこれだ。He の √(2/n)はその 2 を戻す。 GPT-2 の幅 768 では標準偏差 0.051 対 0.036 —— 訓練できる深い ReLU 網と消えていく網の差はここだけだ。

07

積を 1 の近くに保つもの

構造的な手当てが 2 つと、力ずくのものが 1 つ。合わせて、100 層のネットワークが訓練できてしまう理由になる。

残差ブロックは f(x) ではなく x + f(x) を計算するので、その局所微分は f′(x) ではなく 1 + f′(x) になる。この 1は、勾配が通れば何も掛けられずに済む経路だ。

24 層を 2 通りに描いてある。素の鎖と残差の鎖、そして恒等パスをちょうど 1 の位置に横線で引いてある。分岐自身の傾きを動かしてほしい:

分岐の傾き 0.05 —— 素 6.0e−32、残差 3.2e+0

初期の傾き 0.05 —— 弱い分岐で、うまく初期化されたブロックはこう見える —— では、素の鎖は 6.0e−32 に、残差の鎖は 3.2e+0 に届く。ただし残差の鎖がどちら側に外れるかに注目してほしい。分岐を まで上げると 1.7e+4 になる。残差は積を直したのではなく、破綻の仕方を消失から増大へ変えただけだ —— そして増大の方は、正規化で押さえられる。

ブロックの間に挟まる LayerNorm の仕事がそれだ。各ブロックの出力を決まった分散に戻すので、いくつ積んでも残差ストリームは無限には膨らまない。同じ発想で出力射影の初期化を 1/√(2N) だけ小さくする —— あとから直すのではなく、最初から f′ を小さくしておく。

それでも運の悪いバッチで全部が崩れたときのために、鈍器が 1 つある。半径 1 の球と、その外に着地した生の勾配だ。先端をどこへでもドラッグして、オプティマイザが実際に受け取るものを見てほしい:

‖g‖ = 2.61、実際に使うのは 1.00
‖g‖ = 2.61、実際に使うのは 1.00

クリップは切り落としではなくスケール変更だ。球の外側で先端を回すと、実際に使われる勾配は向きを保ち、長さだけを失う。先端をへ入れれば何も起きない。nan になるはずの 1 歩を、偏った更新 1 回で済ませてくれる。‖g‖ ≤ 1.0 は GPT-3 が使った値だ。

あとはループ本体だけ。1 行は誰もが忘れる行だ:

for x, y in loader:
    opt.zero_grad()          # or grads add up
    loss = mse(model(x), y)  # keeps every
    loss.backward()          #   activation
    clip_grad_norm_(p, 1.0)  # the ball above
    opt.step()               # w -= lr * grad

backward() は .grad を上書きせず足し込む —— バッチを複数回に分けられるよう、わざとそうなっている。opt.zero_grad() を落としたとき、各ステップが加える量がこれだ:

ステップ 12、累積中

2 歩目は g₁ + g₂、12 歩目は 12 歩ぶん。向きは正しいので例外は飛ばず損失も下がる —— 狂うのは大きさで、1 時間後に発散する。ば、どの棒も 1 の線に落ちる。

不変条件は §05 の 1 行だ。どの層の勾配も、1 つ後ろの層の勾配に局所微分を 1 つ掛けたもの。安いのは掛け算が共有されるから、2 通りに壊れるのは積が複利だから、レシピが要るのは安全な積の因子が 1 の近くだけだからだ。