LSTM は 隠れ状態 $h$ に加えて記憶セル $c$ を導入することで RNN の勾配消失問題を解決します。

記法

本ドキュメントは以下の記法で統一しています。

記号 意味
$\cdot$ 行列のドット積(行列積)
$\odot$ 要素積(アダマール積、element-wise product)

勾配消失問題

勾配消失は 逆伝播 で問題になります( 順伝播 では問題になりません)。

中間変数 $a$ を定義(活性化関数への入力):

\[a = h_{t-1} \cdot W_h + x_t \cdot W_x + b\] \[h_t = \tanh(a)\]

逆伝播と勾配消失のメカニズム

$L$ の $h_{t-1}$ に関する勾配 $\frac{\partial L}{\partial h_{t-1}}$

\[\frac{\partial L}{\partial h_{t-1}} = \left(\frac{\partial L}{\partial h_t} \odot \left(1 - h_t^2\right) \right) W_h^\top\]

活性化関数に tanh を使用している場合、$\frac{\partial h_t}{\partial a_t} = 1 - \tanh^2(a_t) = 1 - h_t^2$ になります。

$L$ の $h_{t-2}$ に関する勾配:$\frac{\partial L}{\partial h_{t-2}}$

\[\frac{\partial L}{\partial h_{t-2}} = \Bigl( \left( \frac{\partial L}{\partial h_t} \odot (1 - h_t^2) \right) W_h^\top \odot (1-h_{t-1}^2) \Bigr) W_h^\top\]

活性化関数に tanh を使用している場合、$\frac{\partial h_t}{\partial a_t} = 1 - \tanh^2(a_t) = 1 - h_t^2$ になります。

[!NOTE]
️ $W_h^\top$ の $\top$ は転置を表します。

t ステップ遡ると $W_h^\top$ の積が t 回繰り返される

\[\frac{\partial L}{\partial h_1} \propto \frac{\partial L}{\partial h_t} \cdot \underbrace{(W_h)^\top \cdot (W_h)^\top \cdot ...... \cdot (W_h)^\top }_{t 回 の行列積}\]

[!NOTE]
$\propto$ は「比例する」の意味です。他の係数を省略して本質的な部分だけ示しています。

勾配消失 2 つの要因

要因 問題
tanh の値域は $-1 < y < 1$ なので 導関数 $1-h^2$ は $0 < 1 - h_t^2 \le 1$ 掛けるたびに小さくなる
$W_h$ のスペクトル半径 $\rho < 1$(固有値の絶対値 < 1) t 乗で 0 に収束

※ $W_h$ のスペクトル半径 $\rho > 1$ の場合は tanh の縮小効果を上回って勾配爆発が起きることもあります。

固有値と勾配消失・爆発の関係

$\rho(W_h) = \max_i \vert{}\lambda_i\vert{}$(最大固有値の絶対値)をスペクトル半径と呼びます。

条件 $W_h^T$ の挙動 結果
$\rho < 1$(全固有値の絶対値 < 1) $\to 0$ 勾配消失
$\rho = 1$ 安定 安定
$\rho > 1$(ある固有値の絶対値 > 1) $\to \infty$ 勾配爆発

[!NOTE]
消失と爆発が同時に発生する場合もあります(ある成分は 0 へ、別の成分は $\infty$ へ)

固有値の詳細は 固有値 を参照してください。

LSTM による解決策

RNN で導入した(勾配消失問題を持つ)隠れ状態 $h_t$ に加えて記憶セル $c_t$ を導入します。
以下に記憶セル $c_t$ が勾配消失の影響を受けない理由を説明します。

記憶セルとゲート $f, i, o$ と候補セル状態 $g$

記号 名称 活性化関数 値の範囲 主な役割
$f$(Forget) 忘却ゲート シグモイド ($\sigma$) $(0, 1)$ 過去の記憶セル $c_{t-1}$ をどれだけ残すか(または捨てるか)を調整
$i$(Input) 入力ゲート シグモイド ($\sigma$) $(0, 1)$ 新しい情報 $g$ をどれだけ記憶セルに追加するかを制御
$g$(Gate Candidate) 候補セル状態 $\tanh$ $(-1, 1)$ 今回の入力 $x_t$ と直前隠れ状態 $h_{t-1}$ から作られる追加情報の候補本体
$o$(Output) 出力ゲート シグモイド ($\sigma$) $(0, 1)$ 更新されたセル状態 $c_t$ のうち、どれだけを外部出力 $h_t$ に渡すかを調整

記憶セル $c_t$ の順伝播

\[c_t = \underbrace{f_t \odot c_{t-1}}_{\text{過去を保持}} + \underbrace{g_t \odot i_t}_{\text{新情報を追加}}\]

記憶セル $c_t$ の逆伝播

$c_t$ を $c_{t-1}$ で微分

\[\frac{\partial c_t}{\partial c_{t-1}} = f_t\]
  • $f_t \odot c_{t-1}$ を $c_{t-1}$ で微分 : $f_t$ が残る
  • $i_t \odot g_t$ を $c_{t-1}$ で微分 : $c_{t-1}$ を含まないので 0 になる(加算項は消える)。つまり記憶セルに影響を与えない

t 時刻について展開

\[\frac{\partial L}{\partial c_0} = \frac{\partial L}{\partial c_T} \cdot \prod_{k=1}^{T} f_k\]

単純 RNN との比較

  単純 RNN LSTM
逆伝播で掛かるもの 固定の $W_h$(変えられない) 時刻 t ごとに異なる $f_t$(学習で制御できる)
値の範囲 固有値に依存(正確にはスペクトル半径) sigmoid なので $(0, 1)$
制御 困難 $f_t \approx 1$ に学習できる
長期依存 学習困難 学習できる

LSTM が勾配消失に強い理由

1. $f_t$ が時刻ごとに独立(異なる値)

  • 単純 RNN は 同じ $W_h$ の t 乗:固有値(スペクトル半径)に支配される
  • LSTM は 毎時刻異なる $f_t$ の積:固有値問題が起きない

2. $f_t$ を学習で制御できる

$f_t$ は $\text{sigmoid}$ なので値の範囲は $(0, 1)$ です。

[!NOTE]

  • $f_t \approx 1$ に学習:勾配がほぼそのまま過去へ流れます(長期記憶を保持)
  • $f_t \approx 0$ に学習:意図的に過去の情報を遮断できます(忘却)

[!NOTE] この仕組みを定常誤差カルーセル(CEC: Constant Error Carousel)と呼ぶ

$f_t$ を制御する方法
\[y = \sigma(x) = \frac{1}{1 + \exp(-x)}\]

\[\begin{aligned} a_t^f = h_{t-1} \cdot W_h^f + x \cdot W_x^f + b \\ f_t = \sigma (a_t^f) \end{aligned}\]

sigmoid 関数( $\sigma$ )に渡す $a_t^f$ を大きくすれば $f_t$ は 1 に近づきます。
もっとも簡単な方法はバイアス $b$ を大きくすることです。

加算構造が勾配のハイウェイとして機能する理由

$c_t$ の更新式を 2 項の加算として見ると:

\[c_t = \underbrace{f_t \odot c_{t-1}}_{\text{項 A:過去の記憶}} + \underbrace{g_t \odot i_t}_{\text{項 B:新情報}}\]

これを直前の記憶セル $c_{t-1}$ で偏微分すると、次のように計算されます。

\[\frac{\partial c_t}{\partial c_{t-1}} = \frac{\partial (f_t \odot c_{t-1})}{\partial c_{t-1}} + \underbrace{\frac{\partial (g_t \odot i_t)}{\partial c_{t-1}}}_{= 0} = f_t\]

ここで重要なポイントは 項 B(新情報)の微分が 0 になって消える 点です。

  • 項 B の微分が 0 になる意味: 新しい情報の追加処理が、$c_{t-1}$(過去の記憶)の勾配経路に対して直接的なノイズや減衰の干渉を与えない(独立している)ことを表します。
  • 乗算モデルとの対比: 仮に $c_t = f_t \odot c_{t-1} \odot (g_t \odot i_t)$ のような乗算更新だった場合、偏微分に $(g_t \odot i_t)$ が毎回掛け合わされ、過去への勾配を急速に潰してしまいます。
  • ハイウェイの維持: 項 B が 0 として評価されるおかげで、項 A 由来の勾配 $f_t$ だけが濁らずに過去へ伝播します。

逆伝播で上流から勾配 $\frac{\partial L}{\partial c_t}$ が流れてきた際、加算ノードはそれをそのまま 1 倍で項 A・項 B の両方に分配します。

\[\frac{\partial L}{\partial (f_t \odot c_{t-1})} = \frac{\partial L}{\partial c_t} \times 1 = \frac{\partial L}{\partial c_t}\]

[!NOTE]
項 A 内部($f_t \odot c_{t-1}$)は乗算のため、$c_{t-1}$ まで遡ると $f_t$ が掛かります($\frac{\partial L}{\partial c_{t-1}} = \frac{\partial L}{\partial c_t} \odot f_t$)。
したがって、最終的な勾配消失の制御は $f_t \approx 1$ に学習・保持させる議論(理由 1・2)に帰着します。
加算構造は「余計な乗算項を遮断して勾配の通り道を確保する」補強であり、両者が合わさって初めて勾配消失の克服が完結します。

LSTM 詳細

LSTM とは短期記憶( short term memory )を長い( long )時間継続すること意味します。 具体的には勾配消失を受けにくい(制御できる)記憶セル $c_t$ を導入します。

順伝播

LSTM 2

LSTM 3

LSTM は状態を保持する隠れ状態 $h_t$ に加えて記憶を表すセル状態 $c_t$ を持ちます。
$c_t$, $h_t$ の順伝播は以下のとおりです。

\[\begin{aligned} c_t = f_t \odot c_{t-1} + g_t \odot i_t \quad \text{※ c の添字が t-1 に注意} \\ h_t = o \odot \tanh(c_t) \quad \text{※ の添字が t に注意} \\ \\ o_t = \sigma( h_{t-1} \cdot W_h^o + x_t \cdot W_x^o + b) \\ f_t = \sigma( h_{t-1} \cdot W_h^f + x_t \cdot W_x^f + b) \\ i_t = \sigma( h_{t-1} \cdot W_h^i + x_t \cdot W_x^i + b) \\ g_t = \tanh( h_{t-1} \cdot W_h^g + x_t \cdot W_x^g + b) \\ \\ W_x^f \neq W_x^i \neq W_x^g \neq W_x^o​ \\ W_h^f \neq W_h^i \neq W_h^g \neq W_h^o \end{aligned}\]
  • $f_t \odot c_{t-1}$:過去の記憶を $f_t$ の割合で残す
  • $g_t \odot i$:新しい情報 $g_t$ を $i_t$ の割合で追加する
  • $\tanh(c_t)$:セル状態を $(-1, 1)$ に正規化
  • $o \odot$:出力ゲートで「どれだけ外に出すか」を制御
  • $h_t$, $c_t$ は ゲート $f_t, i_t, g_t, o_t$ の影響を受ける
  • $f, i, o$ は情報を反映する割合を表すので活性化関数は sigmoid 関数を使用するので各要素が 0 < y < 1 の $(N, H)$ 次元の配列
  • $g$ は追加する情報の大きさを表すので活性化関数として tanh 関数を使用するので各要素が -1 < y < 1 の $(N, H)$ 次元の配列

効率的に計算するため $f, i, g, o$ の Affine 変換をまとめて計算できるようにします。

\[\begin{aligned} A = h_{t-1} \cdot W_h + x_t \cdot W_x + b \\ \\ W_x^f \neq W_x^i \neq W_x^g \neq W_x^o​ \\ W_h^f \neq W_h^i \neq W_h^g \neq W_h^o \end{aligned}\]

コード

以下 forget を $f$ 、input を $i$、 generate を $g$ 、output を $o$ で表します。

  • バッチサイズ:N
  • 入力の分散表現の次元:D
    • $x_t$ の次元:$(N, D)$
    • $W_x$ の次元: $(D, 4H)$
      • $W_x^f \neq W_x^i \neq W_x^g \neq W_x^o​$ をまとめて $W_x$ と表記(各重み行列は $(D, H)$ )
  • 隠れ状態の分散表現の次元:H
    • $h_t$ の次元: $(N, H)$
    • $W_h$ の次元: $(H, 4H)$
      • $W_h^f \neq W_h^i \neq W_h^g \neq W_h^o$ をまとめて $W_h$ と表記(各重み行列は $(H, H)$ )
    • forget, input, output ゲートおよび generate をまとめて考える場合
  • セル状態(記憶)の次元: H
    • $c_t$ の次元; $(N, H)$
  • $b$ の次元:$(1, 4H)$
import numpy as np

# h_t-1(h_prev), c_t-1(c_prev) は所与とする

input_size = 3
H = hidden_size = 4
bach_size = 10

# 作成  = (4, 4 * 4) 版
W_h = np.random.randn(hidden_size, 4 * hidden_size)  # shape: (4, 16)
# W_x の同様に求める
W_x = np.random.randn(input_size, 4 * hidden_size) # shape: (3, 16)

# Affine 部分( h_prev は h_{t-1} で (3, 4) のミニバッチ
# x @ Wx (10, 16), h_prev @ Wh (10, 16) なので加算可能
A = np.dot(x, Wx) + np.dot(h_prev, Wh) + b

f = A[:, :H]
g = A[:, H:2*H]
i = A[:, 2*H:3*H]
o = A[:, 3*H:]

# 活性化関数に適用
# sigmoid は定義済みとする
f = sigmoid(f)
g = np.tanh(g)
i = sigmoid(i)
o = sigmoid(o)


# c_t, h_t を取得
c_next = f * c_prev + g * i
h_next = o * np.tanh(c_next)

# 補足:取り出し(列方向)
# W_h_f = W_h[:, 0 * hidden_size : 1 * hidden_size]   # forget gate
# W_h_i = W_h[:, 1 * hidden_size : 2 * hidden_size]   # input gate
# W_h_g = W_h[:, 2 * hidden_size : 3 * hidden_size]   # cell gate
# W_h_o = W_h[:, 3 * hidden_size : 4 * hidden_size]   # output gate