深層学習

リカレントニューラルネットワーク

系列を1ステップずつ読み、過去を要約する隠れ状態を次の時刻へ渡すネットワーク。時間方向のパラメータ共有で可変長系列を扱える一方、BPTTでは長期依存の勾配が消失・爆発します。

  • A|中核
  • 深層学習

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

ひとことで言うと

リカレントニューラルネットワーク(RNN)は、入力系列を時刻ごとに読み、過去の情報を隠れ状態に要約して次の時刻へ渡すネットワークです。系列全体を固定長の入力に詰め込まず、同じ状態更新を繰り返すため、系列長が変わっても同じモデルを適用できます。

隠れ状態は、作業記録を毎時点で一枚の引き継ぎメモにまとめる仕組みです。次の担当者は「今回届いた情報」と「前のメモ」だけを受け取り、新しいメモを作ります。メモは過去全体そのものではなく、タスクに必要な情報を残した要約です。

なぜ必要か

固定長のフィードフォワードネットワークで系列を扱うと、長さの違う系列や、同じ情報が現れる位置の違いを別々に処理しがちです。RNNは、時刻の位置ごとに別の重みを持たせる代わりに、同じ遷移規則を全時刻で共有します。これにより、未学習の長さにも同じ規則を適用でき、位置をまたいで統計的な情報を共有できます。

ただし、隠れ状態は任意長の過去系列を固定長ベクトルへ写すため、一般には情報を失う要約です。短期の文脈を次の予測へ渡す用途には自然ですが、遠い時刻の情報を正確に保持する課題では、学習時の勾配が状態をさかのぼれない問題が現れます。

系列処理の方式系列長・位置への対応代償
固定長のフィードフォワード長さや位置ごとの設計に依存可変長・位置の共有が難しい
素の RNN同じ更新規則を全時刻で共有長期依存では勾配が不安定

仕組み

入力を xt\mathbf{x}_t、時刻 tt の隠れ状態を ht\mathbf{h}_t、出力を yt\mathbf{y}_t とします。Wx\mathbf{W}_x は入力から状態への重み、Wh\mathbf{W}_h は前の状態から現在の状態への重み、Wy\mathbf{W}_y は状態から出力への重み、bh,by\mathbf{b}_h,\mathbf{b}_y はバイアス、ϕ\phi は活性化関数です。素の RNN の一例は次です。

ht=ϕ(Whht−1+Wxxt+bh),yt=Wyht+by\mathbf{h}_t = \phi(\mathbf{W}_h\mathbf{h}_{t-1}+\mathbf{W}_x\mathbf{x}_t+\mathbf{b}_h), \qquad \mathbf{y}_t = \mathbf{W}_y\mathbf{h}_t+\mathbf{b}_y

ht−1\mathbf{h}_{t-1} が過去の系列を要約した状態、xt\mathbf{x}_t が現在の入力です。初期状態 h0\mathbf{h}_0 を決めて順に計算すれば、各時刻の状態と出力が得られます。実装上は1つのセルをループで呼びますが、学習時は時間軸に展開した深い計算グラフとして扱います。

ここで重要なのは、時刻ごとに Wx,Wh,Wy\mathbf{W}_x,\mathbf{W}_h,\mathbf{W}_y を複製して別パラメータにするのではなく、同じ値を共有して使うことです。したがって損失 LL の Wh\mathbf{W}_h に対する勾配は、各時刻の寄与の総和になります。

∂L∂Wh=∑t=1T∂L∂Wh∣t\frac{\partial L}{\partial \mathbf{W}_h}=\sum_{t=1}^{T}\frac{\partial L}{\partial \mathbf{W}_h}\bigg|_t

TT は系列の時間ステップ数、右辺の各項は時刻 tt の計算経路を通った寄与です。時間方向に逆向きへ勾配を流す手続きが BPTT(backpropagation through time)です。

状態の経路を kk ステップさかのぼると、概念的には次の積が現れます。

∂ht∂ht−k=∏j=t−k+1t∂hj∂hj−1\frac{\partial \mathbf{h}_t}{\partial \mathbf{h}_{t-k}} =\prod_{j=t-k+1}^{t}\frac{\partial \mathbf{h}_j}{\partial \mathbf{h}_{j-1}}

∂hj/∂hj−1\partial\mathbf{h}_j/\partial\mathbf{h}_{j-1} は1ステップのヤコビ行列です。各行列の影響が時間ステップ数ぶん掛け合わされるため、ノルムが1未満の影響が続けば勾配は急速に小さくなり、1を超える影響が続けば大きくなります。これが、遠い過去へ学習信号が届かない勾配消失と、更新が不安定になる勾配爆発の構造です。爆発には勾配クリッピングが使われますが、消失をそれだけで解決するものではありません。長期の情報経路を保ちやすくするため、後続の RNN ではゲート機構などが導入されました。

系列の入出力は、状態更新と出力をどの時刻で読むかで整理できます。

型入力と出力典型的な読み方
多対一系列を読み、最後または集約した状態から1出力系列全体の分類・判定
多対多(同期)各 xt\mathbf{x}_t に対して各 yt\mathbf{y}_t を出力時刻ごとのラベル付け
多対多(非同期)入力系列を読み終えてから、別の系列を生成エンコーダ・デコーダによる系列変換

エンコーダ・デコーダでは、1つの RNN が入力記号列を固定長ベクトルへ符号化し、別の RNN がその表現から出力記号列を復号します。入力と出力の長さが一致しない系列変換を、多対多の別パターンとして扱える点が要所です。

試験でどう問われるか

問われ方正解に寄る条件引っかけ
状態更新式の記号xt\mathbf{x}_t は現在入力、ht−1\mathbf{h}_{t-1} は前時刻状態、重みは入力経路と再帰経路で分かれるht\mathbf{h}_t を毎回ゼロに戻す/過去状態を入力に使わない
パラメータ数と系列長同じ重みを時間方向に共有するため、系列長に比例して新しい重みは増えない時刻ごとに別の重みを学習する
BPTT の勾配展開した計算グラフを時間方向に逆伝播し、共有パラメータへの寄与を足し合わせる時刻ごとの勾配を別パラメータへ更新する
勾配消失・爆発1ステップのヤコビ行列が時間ステップ数ぶん積になる原因を系列長ではなく、単に学習率だけに帰す
入出力パターン多対一は系列から1出力、多対多は時刻ごとまたは別系列を出力すべての RNN が入力と同じ長さの出力を返す

実装で確かめる

同じ Wh\mathbf{W}_h をループの全時刻で使い、状態が入力を受けて更新される最小例です。時刻ごとに重みを作っていないことが、コード上でも確認できます。

import numpy as np

rng = np.random.default_rng(0)
Wx = rng.normal(size=(2, 3))
Wh = rng.normal(size=(2, 2))
bh = np.zeros(2)
x = rng.normal(size=(4, 3))
h = np.zeros(2)
states = []
for xt in x:
    h = np.tanh(Wx @ xt + Wh @ h + bh)
    states.append(h.copy())
states = np.array(states)
print(states.shape, Wh.shape)

実行結果は (4, 2) (2, 2) です。4時刻ぶんの状態を得ても、再帰に使う重み Wh\mathbf{W}_h の形は1つのままです。学習時にはこの4回の更新を展開し、共有された Wh\mathbf{W}_h に各時刻の勾配を集約します。

取り違えやすいもの

用語素の RNN との切り分け
フィードフォワードネットワーク入力から出力へ循環せず、前時刻の状態を次時刻へ渡さない。固定長入力向けの構成になりやすい
BPTTRNNそのものではなく、時間方向に展開したRNNを訓練するための逆伝播手続き
RNNエンコーダ・デコーダRNNを2つ使い、入力系列を固定長表現にして別系列を生成する構成。素のセルの状態更新とは別の組み合わせ方
ゲート付き RNN素の RNN の長期依存の弱点に対し、情報の保持・更新を制御する機構を加えたもの。ここで扱う式にはゲートはない
注意機構固定長状態だけに全過去を押し込む制約を別の参照方法で扱う機構。RNNの再帰更新そのものではない

想起チェック

RNNが系列長の違いを扱える理由は何か

入力サイズを時刻ごとの遷移として定義し、同じパラメータを時間方向に共有して使うためです。

BPTTで勾配消失・爆発が起きる式の構造は何か

時間をさかのぼる勾配に、1ステップのヤコビ行列がステップ数ぶん掛け合わされます。積のノルムが小さくなり続ければ消失し、大きくなり続ければ爆発します。

多対一とエンコーダ・デコーダの違いは何か

多対一は入力系列から1つの出力を読む構成です。エンコーダ・デコーダは入力系列を固定長表現へ符号化し、そこから別の出力系列を復号します。

素のRNNの長期依存の限界に対して、後続手法は何を導入したか

情報の保持や更新を制御するゲート機構などです。ここでは素のRNNの状態更新と、限界が生じる勾配の積までを対象にします。

出典