本文へスキップ
ひもとくAI

勾配消失と LSTM・GRU:ゲートで覚える・忘れるを調節する

レッスン 3/7

RNN の勾配消失と勾配爆発

さかのぼるほど、勾配が小さくなる(大きくなる)

BPTT では、出力の誤差を、時刻をさかのぼって伝えていきます。1つさかのぼるたびに、同じ重み WhW_h と、tanh⁡\tanh の傾き(1以下)が掛けられます。同じような数を何十回も掛けると、1より小さければ0に、1より大きければ爆発的に大きくなります(0.5を10回掛けると約0.001、2を10回掛けると1024)。

NumPy で作った RNN(隠れ状態16個、系列の長さ50)で、重み WhW_h の大きさ(固有値の大きさの最大値。第11章)を変えて、何ステップさかのぼったところに、どれくらいの大きさの勾配が届くかを測りました。

重み WhW_h の大きさ1つ前5つ前10個前20個前49個前
0.56.3×10−36.3 \times 10^{-3}1.7×10−41.7 \times 10^{-4}2.1×10−62.1 \times 10^{-6}2.3×10−102.3 \times 10^{-10}2.9×10−212.9 \times 10^{-21}
1.01.4×10−21.4 \times 10^{-2}4.3×10−34.3 \times 10^{-3}1.1×10−31.1 \times 10^{-3}6.1×10−56.1 \times 10^{-5}5.6×10−85.6 \times 10^{-8}
1.52.4×10−22.4 \times 10^{-2}2.2×10−22.2 \times 10^{-2}1.8×10−21.8 \times 10^{-2}1.3×10−21.3 \times 10^{-2}5.7×10−35.7 \times 10^{-3}
3.04.6×10−24.6 \times 10^{-2}1.1×10−11.1 \times 10^{-1}4.3×10−14.3 \times 10^{-1}7.87.85.3×1045.3 \times 10^{4}
  • 重みが小さいと(0.5)、49ステップ前に届く勾配は 10−2110^{-21} で、ほぼ0です(勾配消失)。遠い過去の入力を、学習で調整できません
  • 重みが大きいと(3.0)、さかのぼるほど勾配が大きくなり、5×1045 \times 10^{4} にもなりました(勾配爆発)。重みの更新が大きすぎて、学習が壊れてしまいます

第8章では、層を深く重ねたときの勾配消失を学びました。RNN は、時間方向に広げると「系列の長さだけ層がある、とても深いネットワーク」なので、同じ問題が起こるのです。

勾配クリッピング:爆発を防ぐ

勾配爆発には、勾配の大きさ(全体の長さ)が決まった値を超えたら、向きはそのままで、長さをその値まで縮める方法がよく使われます。これを 勾配クリッピング と呼びます。

g←g×上限∣g∣(∣g∣>上限 のとき)g \leftarrow g \times \frac{\text{上限}}{|g|} \quad (|g| > \text{上限} \text{ のとき})

PyTorch では、nn.utils.clip_grad_norm_(model.parameters(), 1.0) の1行で使えます。一方、勾配消失は、クリッピングでは防げません。

実験:最初の値を、最後まで覚えておけるか

最初の時刻の値が +2 か −2 かを、雑音が続いた後の系列の最後に答える課題を作り、NumPy の RNN(Adam、勾配の各値を−1〜1に切りそろえる)で学習させました(テストの正解率、乱数の種を3通り)。

系列の長さ2〜204060100
正解率ほぼ1.01.0・0.494・1.00.506・0.544・0.5020.512・0.494・1.0

長さ60では、3回とも当て推量(0.5)と同じでした。系列が長くなると、最初の入力まで勾配が届きにくくなり、覚えることを学習できなくなったと考えられます(長さ100で1回成功したのは、乱数の種による偶然と考えられます)。

LSTM と GRU:ゲートで、覚える・忘れるを調節する

LSTM:記憶の通り道と、3つの蛇口

LSTM(Long Short-Term Memory、長短期記憶)は、RNN の隠れ状態とは別に、セル cc という「記憶の通り道」を持ちます。そして、情報の出し入れを、ゲート と呼ばれる3つの蛇口で調節します。ゲートは、シグモイド関数(第5章、0〜1の値)の出力で、「どれだけ通すか」の割合です。

ゲート役割0に近いと1に近いと
忘却ゲート ff前のセルの記憶を、どれだけ残すか忘れる残す
入力ゲート ii新しい情報を、どれだけセルに書き込むか書き込まない書き込む
出力ゲート ooセルの記憶を、どれだけ隠れ状態として外に出すか出さない出す
ct=f⊙ct−1+i⊙c~t,ht=o⊙tanh⁡(ct)c_t = f \odot c_{t-1} + i \odot \tilde{c}_t, \qquad h_t = o \odot \tanh(c_t)

c~t\tilde{c}_t は、今の入力と前の隠れ状態から作った「新しい情報の候補」です。⊙\odot は、第11章の アダマール積(同じ位置どうしの掛け算)で、ゲートの値を、数ごとに掛けます。

大切なのは、セルの更新が 掛け算の繰り返しではなく、足し算 になっていることです。忘却ゲートが1に近く、入力ゲートが0に近ければ、ct≈ct−1c_t \approx c_{t-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通り、勾配クリッピングあり)。

系列の長さRNN の二乗誤差LSTM の二乗誤差
500.128・0.0770.010・0.031
1000.162・0.1610.025・0.159

LSTM は、長さ50では2回とも、長さ100でも1回は、離れた2か所の数を覚えて足すことを学習しました。RNN は、長さ100では、いつも1と答えるのとほとんど変わりませんでした。

ただし、LSTM がいつも勝つわけではない

前のスライドの「最初の値を覚えておく」課題(値を ±1 にして、少し難しくしたもの)でも比べると、結果は逆でした。

系列の長さRNN の正解率LSTM の正解率LSTM(忘却ゲートのバイアスの初期値を1)の正解率
500.548・1.0・0.970.498・0.5・0.51.0・0.496・1.0
1000.96・1.0・0.4920.49・0.534・0.5080.494・0.53・0.458

この課題では、勾配クリッピングを使った 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つにした軽い版
  • どのモデルがよいかは、課題と設定で変わる。実験で確かめる

演習

演習を読み込んでいます…