深層学習

LSTM

長期依存を学習するため、セル状態への加算経路と入力・出力ゲートを持たせ、時間方向の誤差を保持しながら記憶の書き込みと読み出しを制御するRNNです。

  • A|中核
  • 深層学習

数式の記号で止まったら 記号の読み方 (∂・⊙・転置・上付き添字を、読み方から)

ひとことで言うと

LSTM(Long Short-Term Memory)は、系列を1時刻ずつ処理しながら、短期的な出力とは別に長く保持するセル状態 ct\mathbf{c}_t を持つRNNです。セル状態を毎時刻の非線形変換で作り直さず、ゲートで選んだ量を加算して更新することで、長い時間を隔てた情報と誤差の通り道を確保します。

セル状態は、作業机の上に置いたメモです。入力ゲートはメモへ何を書くか、忘却ゲートは何を消すか、出力ゲートはメモのどの部分を今の返答に見せるかを決めます。机そのものを毎回書き直すのではなく、必要な箇所だけ更新するので、数十ステップ前の情報を残したまま現在の判断もできます。

なぜ必要か

経路通常のRNNLSTM
記憶の更新非線形変換で隠れ状態を作り直すセル状態へ保持と追加を加算する
誤差の通り道時刻ごとの変換の積セル状態の加算経路を使える

通常のRNNでは、隠れ状態を ht=tanh⁡(Wht−1+Uxt+b)\mathbf{h}_t = \tanh(W\mathbf{h}_{t-1}+U\mathbf{x}_t+\mathbf{b}) のように毎時刻変換します。損失を過去へ戻すと、同じ形のヤコビアンが何度も掛かります。各時刻の微分の大きさが1未満なら勾配は急速に小さくなり、1を超える方向があれば爆発します。そのため、現在の出力に必要な情報がかなり前の入力に依存していても、学習信号がそこまで届きません。

LSTMが解こうとしたのは、この「長期依存を表現できない」問題と、その原因である誤差の消失・爆発です。核心は、記憶を大きな非線形変換の中へ押し込めるのではなく、セル状態に専用の経路を与えたことです。セル状態の更新を概念的に ct=ct−1+追加−削除\mathbf{c}_t=\mathbf{c}_{t-1}+\text{追加}-\text{削除} と見ると、時間をまたぐ微分に恒等写像に近い経路が現れます。入力を毎回保存するのではなく、保存すべき時だけ書き込ませるので、長期情報と新しい情報が同じ場所で無秩序に混ざりません。

ここでいう「誤差が保たれる」は、どんな長さの系列でも完全に勾配が一定になるという意味ではありません。忘却ゲートなどの値、損失、パラメータ、入力によって勾配は変わります。重要なのは、セル状態の加算経路が、通常のRNNのように毎時刻必ず非線形関数を通る構造ではないことです。CEC(constant error carousel)はこの設計意図を指す呼び名で、ゲートはその経路へいつ書き込み、いつ読み出すかを学習可能にします。

仕組み

入力 xt\mathbf{x}_t、前時刻の隠れ状態 ht−1\mathbf{h}_{t-1}、セル状態 ct−1\mathbf{c}_{t-1} を結合したベクトルを zt=[ht−1,xt]\mathbf{z}_t=[\mathbf{h}_{t-1},\mathbf{x}_t] とします。標準的な表記の一例は次です。

ft=σ(Wfzt+bf),it=σ(Wizt+bi),c~t=tanh⁡(Wczt+bc)\mathbf{f}_t=\sigma(W_f\mathbf{z}_t+\mathbf{b}_f),\quad \mathbf{i}_t=\sigma(W_i\mathbf{z}_t+\mathbf{b}_i),\quad \tilde{\mathbf{c}}_t=\tanh(W_c\mathbf{z}_t+\mathbf{b}_c) ct=ft⊙ct−1+it⊙c~t\mathbf{c}_t=\mathbf{f}_t\odot\mathbf{c}_{t-1}+\mathbf{i}_t\odot\tilde{\mathbf{c}}_t ot=σ(Wozt+bo),ht=ot⊙tanh⁡(ct)\mathbf{o}_t=\sigma(W_o\mathbf{z}_t+\mathbf{b}_o),\qquad \mathbf{h}_t=\mathbf{o}_t\odot\tanh(\mathbf{c}_t)

σ\sigma は0から1へ写すシグモイド関数、tanh⁡\tanh は候補値と出力の非線形変換、⊙\odot は要素ごとの積です。ft\mathbf{f}_t は古いセル状態を残す割合、it\mathbf{i}_t は候補 c~t\tilde{\mathbf{c}}_t を書き込む割合、ot\mathbf{o}_t はセル状態を隠れ状態へ出す割合です。実装では3つのゲートと候補を別々に線形変換する代わりに、4つ分を一度の行列積で計算して分割することが多いですが、意味は変わりません。

試験やデバッグで最も大切なのは式を暗記することではなく、経路の役割を分けることです。セル状態には「保持」と「候補の追加」があり、隠れ状態には「セルからの読み出し」があります。忘却ゲートが1に近く入力ゲートが0に近い時、ct\mathbf{c}_t はほぼ ct−1\mathbf{c}_{t-1} のままです。このときセル状態を時間方向に微分すると、主経路には ft\mathbf{f}_t が現れ、ft\mathbf{f}_t が長い区間で1に近ければ勾配は保たれます。一方、隠れ状態 ht\mathbf{h}_t は出力ゲートと tanh⁡\tanh を通るため、常にそのまま伝わるわけではありません。CECが守るのはセル状態側の長い経路です。

試験でどう問われるか

問われ方正解に寄る条件引っかけ
長期依存と勾配消失の関係セル状態の加算経路が、時間方向に誤差を運ぶ「ゲートがあるから必ず勾配は一定」
セル状態の更新式古い状態の保持と候補の書き込みを要素積で混ぜる候補を常に全量上書きする
入力・出力ゲートの役割入力は書き込み、出力は隠れ状態への読み出しを制御する出力ゲートがセル状態を更新する
CECの意味加算型のセル状態経路により定数誤差を狙う設計CECを活性化関数や損失関数と扱う
実装のテンソル確認ゲートの次元は隠れ次元、セル状態と要素積できるバッチ次元をゲート次元と混同する

実装で確かめる

NumPyで1時刻の更新だけを書き、入力ゲートを閉じた場合に候補がセル状態へ反映されないことを確認します。出力ゲートを変えると見える ht\mathbf{h}_t は変わりますが、ct\mathbf{c}_t は変わらない点が、書き込みと読み出しの分離です。

import numpy as np

x = np.array([0.2, -0.4])
h_prev = np.array([0.1, 0.3])
c_prev = np.array([2.0, -1.0])
sigmoid = lambda a: 1 / (1 + np.exp(-a))
z = np.r_[h_prev, x]
f = sigmoid(np.array([6.0, 6.0]))
i = sigmoid(np.array([-8.0, -8.0]))
g = np.tanh(np.array([1.0, -1.0]))
o = sigmoid(np.array([0.0, 0.0]))
c = f * c_prev + i * g
h = o * np.tanh(c)
print(np.round(c, 3), np.round(h, 3))

このコードでは ft\mathbf{f}_t はほぼ1、it\mathbf{i}_t はほぼ0なので、セル状態はほぼ以前の値を保ちます。実装を拡張するときは、zを使った4つの線形変換、バッチ次元、初期状態のゼロ埋めを順に確認すると、式の転記ミスを見つけやすくなります。

c が長期記憶を持つからといって、h を使わずに予測できるわけではありません。出力ゲートで必要な情報だけを h へ出し、次時刻のゲート計算には h と入力を使います。セル状態をそのまま外へ出す実装は、標準的なLSTMの読み出し制御を失います。

取り違えやすいもの

構造状態更新の見方長期依存との関係
通常のRNN隠れ状態を毎時刻、非線形変換で作り直す時間方向の積で勾配消失・爆発が起きやすい
LSTMセル状態を保持・消去・書き込みし、隠れ状態へ読み出す加算経路で長い誤差の通路を設ける
CECLSTM内部の設計上の発想ゲートの操作対象ではなく、セル状態の長い経路を指す
BPTTRNNを時間方向に展開して逆伝播する計算法LSTMでも使う。LSTMそのものと同義ではない

想起チェック

LSTMが通常のRNNで起きる長期依存の問題に対して導入した中心的な経路は何か

セル状態を毎時刻の非線形変換で上書きせず、保持した状態へ候補を加算する経路です。CECと呼ばれるこの発想により、時間方向に誤差が通る経路を確保します。

入力ゲートと出力ゲートは、それぞれ何を制御するか

入力ゲートは候補をセル状態へどれだけ書き込むか、出力ゲートはセル状態を隠れ状態へどれだけ読み出すかを制御します。書き込みと読み出しは別の操作です。

忘却ゲートが1に近く入力ゲートが0に近いとき、セル状態はどうなるか

ct\mathbf{c}_t はほぼ ct−1\mathbf{c}_{t-1} のままです。古い記憶を保持する設定であり、候補が計算されてもセル状態への書き込みはほとんど起きません。

CECがあるなら、LSTMの勾配は必ず消失しないか

必ずではありません。セル状態の主経路を保ちやすくする設計ですが、ゲート値、出力側の非線形変換、損失やパラメータによって勾配は変化します。

出典