RNN の勾配消失と勾配爆発
さかのぼるほど、勾配が小さくなる(大きくなる)
BPTT では、出力の誤差を、時刻をさかのぼって伝えていきます。1つさかのぼるたびに、同じ重み Wh と、tanh の傾き(1以下)が掛けられます。同じような数を何十回も掛けると、1より小さければ0に、1より大きければ爆発的に大きくなります(0.5を10回掛けると約0.001、2を10回掛けると1024)。
NumPy で作った RNN(隠れ状態16個、系列の長さ50)で、重み Wh の大きさ(固有値の大きさの最大値。第11章)を変えて、何ステップさかのぼったところに、どれくらいの大きさの勾配が届くかを測りました。
- 重みが小さいと(0.5)、49ステップ前に届く勾配は 10−21 で、ほぼ0です(勾配消失)。遠い過去の入力を、学習で調整できません
- 重みが大きいと(3.0)、さかのぼるほど勾配が大きくなり、5×104 にもなりました(勾配爆発)。重みの更新が大きすぎて、学習が壊れてしまいます
第8章では、層を深く重ねたときの勾配消失を学びました。RNN は、時間方向に広げると「系列の長さだけ層がある、とても深いネットワーク」なので、同じ問題が起こるのです。
勾配クリッピング:爆発を防ぐ
勾配爆発には、勾配の大きさ(全体の長さ)が決まった値を超えたら、向きはそのままで、長さをその値まで縮める方法がよく使われます。これを 勾配クリッピング と呼びます。
g←g×∣g∣上限(∣g∣>上限 のとき)
PyTorch では、nn.utils.clip_grad_norm_(model.parameters(), 1.0) の1行で使えます。一方、勾配消失は、クリッピングでは防げません。
実験:最初の値を、最後まで覚えておけるか
最初の時刻の値が +2 か −2 かを、雑音が続いた後の系列の最後に答える課題を作り、NumPy の RNN(Adam、勾配の各値を−1〜1に切りそろえる)で学習させました(テストの正解率、乱数の種を3通り)。
長さ60では、3回とも当て推量(0.5)と同じでした。系列が長くなると、最初の入力まで勾配が届きにくくなり、覚えることを学習できなくなったと考えられます(長さ100で1回成功したのは、乱数の種による偶然と考えられます)。
LSTM と GRU:ゲートで、覚える・忘れるを調節する
LSTM:記憶の通り道と、3つの蛇口
LSTM(Long Short-Term Memory、長短期記憶)は、RNN の隠れ状態とは別に、セル c という「記憶の通り道」を持ちます。そして、情報の出し入れを、ゲート と呼ばれる3つの蛇口で調節します。ゲートは、シグモイド関数(第5章、0〜1の値)の出力で、「どれだけ通すか」の割合です。
ct=f⊙ct−1+i⊙c~t,ht=o⊙tanh(ct)
c~t は、今の入力と前の隠れ状態から作った「新しい情報の候補」です。⊙ は、第11章の アダマール積(同じ位置どうしの掛け算)で、ゲートの値を、数ごとに掛けます。
大切なのは、セルの更新が 掛け算の繰り返しではなく、足し算 になっていることです。忘却ゲートが1に近く、入力ゲートが0に近ければ、ct≈ct−1 で、記憶がそのまま次の時刻に運ばれます。さかのぼる勾配も、この通り道を通れば、忘却ゲートの値を掛けるだけで伝わるので、RNN のように急に消えにくくなります。
GRU:ゲートを2つにした、軽い版
GRU(Gated Recurrent Unit)は、セルを持たず、隠れ状態を直接、2つのゲート(更新ゲート と リセットゲート)で調節します。LSTM より重みが少なく計算が軽いのに、多くの課題で同じくらいの性能を出します。
実験:RNN と LSTM を比べる(PyTorch)
LSTM の学習は重いので、運営側で PyTorch(2.14.0)を使って比べました。足し算の課題(LSTM を提案した論文でも使われた課題)は、各時刻に「0〜1の数」と「印」の2つが入り、印の付いた2か所の数の和を、系列の最後に答えるものです。いつも平均の1と答えると、二乗誤差は約0.167になります(乱数の種を2通り、勾配クリッピングあり)。
LSTM は、長さ50では2回とも、長さ100でも1回は、離れた2か所の数を覚えて足すことを学習しました。RNN は、長さ100では、いつも1と答えるのとほとんど変わりませんでした。
ただし、LSTM がいつも勝つわけではない
前のスライドの「最初の値を覚えておく」課題(値を ±1 にして、少し難しくしたもの)でも比べると、結果は逆でした。
この課題では、勾配クリッピングを使った RNN の方がよく学習し、LSTM は多くの場合で失敗しました。忘却ゲートのバイアスの初期値を1にする(はじめは「忘れにくい」状態から始める、よく使われる工夫)と、長さ50では改善しました。LSTM は長い依存関係を学習しやすい 仕組み を持っていますが、実際に学習できるかは、課題・初期値・学習の設定しだいです。
(PyTorch の結果は、同じ設定で実行し直しても、小数第2〜3位が少し変わることがありました。計算の順番の違いなどによる、ごく小さな誤差が、学習の途中で広がるためと考えられます。第9章の再現性の話と同じ種類の問題です)
PyTorch で LSTM を使う、よくある誤解と振り返り
import torch
from torch import nn
lstm = nn.LSTM(input_size=2, hidden_size=32, batch_first=True) # RNN と同じ書き方
gru = nn.GRU(input_size=2, hidden_size=32, batch_first=True)
out = nn.Linear(32, 1)
X = torch.randn(128, 100, 2) # (データの数, 時刻の数, 入力の数)
hs, (h_last, c_last) = lstm(X) # LSTM は、隠れ状態 h とセル c の両方を返す
pred = out(hs[:, -1])
loss = ((pred - torch.ones(128, 1)) ** 2).mean()
loss.backward()
nn.utils.clip_grad_norm_(lstm.parameters(), 1.0) # 勾配の長さを1までに縮める(勾配クリッピング)
nn.RNN を nn.LSTM や nn.GRU に変えるだけで使えます。勾配クリッピングは、backward() の後、重みを更新する前に行います。
よくある誤解
- 「LSTM を使えば、どんなに長い系列でも覚えられる」: 実験のとおり、長さ100の足し算の課題では、2回中1回は失敗しました。覚えておける長さには限りがあり、学習の設定しだいで失敗もします。さらに長い文脈を扱う方法が、次のレッスンからの Attention です
- 「勾配クリッピングで、勾配消失も防げる」: クリッピングは、大きすぎる勾配を縮めるだけです。小さすぎる勾配を大きくはしません
- 「ゲートは、人が手で開け閉めする」: ゲートの値は、入力と前の隠れ状態から、学習した重みで自動的に計算されます。「いつ覚え、いつ忘れるか」もデータから学習します
振り返り
- RNN は、時刻をさかのぼるたびに同じ重みが掛けられるので、勾配が消えたり爆発したりする。爆発は勾配クリッピングで防ぐ
- LSTM は、セル(記憶の通り道)と、忘却・入力・出力の3つのゲートで、情報の出し入れを調節する。セルの更新が足し算なので、勾配が消えにくい
- GRU は、ゲートを更新・リセットの2つにした軽い版
- どのモデルがよいかは、課題と設定で変わる。実験で確かめる