RNN 逆伝播
ゼロから作るDeep Learning ❷ ―自然言語処理編 斎藤康毅 (著) の読書メモです。
記法
本ドキュメントは以下の記法で統一します。
| 記号 | 意味 |
|---|---|
| $\cdot$ | 行列積(ドット積) |
| $\odot$ | 要素積(アダマール積、element-wise product) |
前提
- corpus size: 1000
- V( vocab_size ):コーパスから重複を排除した単語 ID の数(語彙数)
- ソースコードは 418
- N:ミニバッチサイズ
- ソースコードは 10
- T:系列長(sequence length)
- ソースコードは 5
- E( wordvec_size ):単語の分散表現の次元数(要素数)
- ソースコードは 100
- H( hidden_size ): RNNの隠れ状態の次元数(要素数)
- ソースコードは 100
※ ソースコードの $E$, $H$ はともに 100 ですが異なる値を設定できます。
RNN レイヤ


RNN レイヤの逆伝播
Deep Learning における勾配は損失 $L$ に関する勾配 $\frac{\partial L}{\partial a_t}$ 、 $\frac{\partial L}{\partial W_x}$ 、 $\frac{\partial L}{\partial W_h}$ 、 $\frac{\partial L}{\partial h_{t-1}}$ 、 $\frac{\partial L}{\partial b}$ です。
各勾配は時刻 $T = t$ ごとに求めます。
逆伝播は 2 ステップで考ます。
Step 1. $L$ の $a_t$ に関する勾配 $\frac{\partial L}{\partial a_t}$ を求める(要素積 1 のみ):
\[\frac{\partial L}{\partial a_t} = \frac{\partial L}{\partial h_{t}} \odot \frac{\partial h_t}{\partial a_t} = dh_t \odot (1 - h_t \odot h_t) \quad \text{shape: } (N, H)\]活性化関数に tanh を使用している場合、$\frac{\partial h_t}{\partial a_t} = 1 - \tanh^2(a_t) = 1 - h_t^2$ になります。
| 記号 | 意味 | shape |
|---|---|---|
| $\odot$ | 要素積(element-wise) | — |
| $dh_t$ | 上流からの勾配 | $(N, H)$ |
| $h_t \odot h_t$ | $h_t$ の各要素を2乗 | $(N, H)$ |
| $1$ | 全要素が1の行列 | $(N, H)$ |
$dh_t^{(\text{出力層})}$ の正体(TimeAffine / TimeSoftmaxWithLoss との接続)
上記の $dh_t^{(\text{出力層})}$ は、実際には以下の経路で計算されたものが TimeRNN レイヤに渡ってきています。
-
TimeSoftmaxWithLoss.backward() Softmax に入る前のスコア $s$ に対する勾配 $\dfrac{\partial L}{\partial s} = q - y$ を求める ($q$:モデルの予測確率、$y$:正解の one-hot ベクトル。詳細は RNN 時系列 Softmax-with-Loss レイヤの逆伝播 を参照)
- TimeAffine.backward()
- で得た $(N, T, V)$ の勾配を受け取り、Affine レイヤの逆伝播公式 (「forward で掛けた重み行列を転置して右から掛ける」)に従って 隠れ状態 $h_t$ に関する勾配 $(N, T, H)$ を算出する
- TimeRNN.backward()
- で得た勾配が
dhs(各時刻の出力層からの勾配)として渡され、 ループ内で次時刻からの勾配dh_nextと合算される:
- で得た勾配が
dh = dhs[:, t, :] + dh_next
これが本ドキュメント冒頭の $dh_t^{(\text{出力層})} + dh_t^{(\text{次時刻})}$ に対応します。
つまり、
\[\underbrace{q - y}_{\text{Softmax-with-Loss}} \;\xrightarrow{\text{TimeAffine}}\; dh_t^{(\text{出力層})} \;\xrightarrow{\;+\; dh_t^{(\text{次時刻})}\;}\; \frac{\partial L}{\partial a_t} \;\xrightarrow{\text{RNNセル}}\; \frac{\partial L}{\partial W_x}, \frac{\partial L}{\partial W_h}, \frac{\partial L}{\partial h_{t-1}}, \ldots\]という一本の勾配の流れとして、2つのドキュメントの内容が接続されています。
Step 2. $L$ の $a_t$ に関する勾配 $\frac{\partial L}{\partial a_t}$ を起点に各勾配を求める(ドット積):
| 勾配 | 式 | shape確認 |
|---|---|---|
| $\frac{\partial L}{\partial W_x}$ | $x_{t}^\top \cdot \frac{\partial L}{\partial a_t}$ | $(E,N)(N,H) = (E,H)$ |
| $\frac{\partial L}{\partial W_h}$ | $h_{t-1}^\top \cdot \frac{\partial L}{\partial a_t}$ | $(H,N)(N,H) = (H,H)$ |
| $\frac{\partial L}{\partial x_t}$ | $\frac{\partial L}{\partial a_t} \cdot W_x^\top$ | $(N,H)(H,E) = (N,E)$ |
| $\frac{\partial L}{\partial h_{t-1}}$ | $\frac{\partial L}{\partial a_t} \cdot W_h^\top$ | $(N,H)(H,H) = (N,H)$ |
| $\frac{\partial L}{\partial b}$ | $\sum_{n=1}^{N} \frac{\partial L}{\partial a_t}$ | $(N,H) \to (H,)$ 数学的には $(1,H)$ だがコードの実装上は $(H,)$。以下 $(H,)$ と記載する |
※ 数学的には行列積の定義より $(E, N)(N, H) = (E, H)$ が必ず成り立ちますが、真ん中の $N$ が消える(縮約される)意味をミニバッチ全員分の損失を重みに集約しているからと解釈すると理解が深まります。
パターン:重み行列の勾配は「forward で掛けた相手を転置して左から掛ける」、
入力・隠れ状態の勾配は「forward で掛けた重み行列を転置して右から掛ける」。
$L$ の $a_t$ に関する勾配 $\frac{\partial L}{\partial a_t}$
\[\frac{\partial L}{\partial a_t} = \frac{\partial L}{\partial h_t} \odot \frac{\partial h_t}{\partial a_t} = dh_t \odot (1 - h_t \odot h_t) \quad \text{shape: } (N, H)\]このとき $ \frac{\partial L}{\partial h_t}$ は以下の勾配の和です。
- 各時刻の出力層からの勾配
- 時刻をまたぐ再帰的な勾配
TimeRNN.backward() ではdh = dhs[:, t, :] + dh_nextとして両者を加算しいます。
NOTE
RNN(単一セル)と TimeRNN の役割分担
- RNN.backward() は「上流から渡された $dh_t$」をそのまま使って $\partial L/\partial a_t = dh_t \odot (1-h_t\odot h_t)$ を計算するだけで、 その $dh_t$ の中身(出力層からの勾配なのか、次時刻からの勾配なのか)は関知しない。
- TimeRNN.backward() がループの中で、RNN に渡す前に
dh = dhs[:, t, :] + dh_nextとして2経路分の勾配を先に合算し、 その完成品のdhをlayer.backward(dh)として RNN セルへ渡している。- つまり1つ目の式の $dh_t$ は「TimeRNN によってすでに2経路の和が 済んだ完成品」であり、2つ目の式はその中身を分解して示したもの。 両者は矛盾ではなく、粒度(RNNセル単体 vs TimeRNN全体)の違い。
- なお、この「時刻をまたぐ勾配($dh_t^{(\text{次時刻})}$)」は stateful の True/False に関係なく、1回の Truncated BPTT ブロック ($T$ステップ分)内では常に発生する。stateful が制御するのは forward の隠れ状態をバッチをまたいで引き継ぐかどうかだけであり、 時間方向のBPTT自体とは別の話。
$L$ の $W_x$ に関する勾配 $\frac{\partial L}{\partial W_x}$
\[\frac{\partial L}{\partial W_x} = x_{t}^\top \cdot \frac{\partial L}{\partial a_t}\]- $\frac{\partial a_t}{\partial W_x} = x$ : $a_t = h_{t-1} \cdot W_h + x_{t} \cdot W_x + b$ を $W_x$ で微分
- forward で $x$ を左から掛けていたので、転置して $x_{t}^\top$ を左から掛ける
NOTE
なぜ $x_{t}^\top$ を左から掛けるのかforward では:
$a_t = x \cdot W_x \quad (N,E)(E,H) = (N,H)$
逆伝播の原則は「forward で左にいた行列を転置して左から掛ける」です。
forward で $x$ は $W_x$ の左にいたので $x_{t}^\top$ を左から掛けます。$\frac{\partial L}{\partial W_x} = x_{t}^\top \cdot da \quad (E,N)(N,H) = (E,H)$
$L$ の $W_h$ に関する勾配 $\frac{\partial L}{\partial W_h}$
\[\frac{\partial L}{\partial W_h} = h_{t-1}^\top \cdot \frac{\partial L}{\partial a_t}\]- $\frac{\partial a_t}{\partial W_h} = h_{t-1}$ : $a_t = h_{t-1} \cdot W_h + x \cdot W_x + b$ を $W_h$ で微分
- forward で $h_{t-1}$ を左から掛けていたので、転置して $h_{t-1}^\top$ を左から掛ける
$L$ の $x_t$ に関する勾配 $\frac{\partial L}{\partial x_t}$
\[\frac{\partial L}{\partial x} = \frac{\partial L}{\partial a_t} \cdot W_x^\top\]- $\frac{\partial a_t}{\partial x} = W_x$ : $a_t = h_{t-1} \cdot W_h + x \cdot W_x + b$ を $x$ で微分
- forward で $W_x$ を右から掛けていたので、転置して $W_x^\top$ を右から掛ける
$L$ の $h_{t-1}$ に関する勾配 $\frac{\partial L}{\partial h_{t-1}}$
\[\frac{\partial L}{\partial h_{t-1}} = \frac{\partial L}{\partial a_t} \cdot W_h^\top\]- $\frac{\partial a_t}{\partial h_{t-1}} = W_h$ : $a_t = h_{t-1} \cdot W_h + x \cdot W_x + b$ を $h_{t-1}$ で微分
- forward で $W_h$ を右から掛けていたので、転置して $W_h^\top$ を右から掛ける
- この勾配が1つ前の時刻 $t-1$ の RNN へ伝播することで BPTT が実現されます。
$L$ の $b$ に関する勾配 $\frac{\partial L}{\partial b}$
\[a_t = h_{t-1}W_h + x_{t}W_x + b \quad \text{※} \ b = \left( b_1, \ldots, b_H \right)\]バイアスは各時点 $t$ ごとに考えるのではなく、ミニバッチ・全時刻を通じて共有される ただ1つのパラメータとして考えます。
- $\frac{\partial a_{t,n}}{\partial b} :すべの要素が 1 のベクトル \left(1,\ldots, 1 \right)(以降 $\mathbf{1}_H$ と表記するこもとまります。
- $b$ はミニバッチ $N$ 個・全時刻 $T$ 個すべてに同一の値として加算されるため、逆伝播では $N$ 方向・$T$ 方向の両方に合算する (実装上は $N$ 方向はドット積内のsumで、$T$ 方向はTimeRNNのループでの加算で実現される)
BPTT(Backpropagation Through Time)
以下の RNN を時間方向に展開したネットワーク全体に対する逆伝播アルゴリズムを BPTT(Backpropagation Through Time)と呼びます。
各時刻の勾配が2経路の和になる点は前述の通りです ($L$ の $a_t$ に関する勾配 の NOTE を参照)。
Truncated BPTT
実装では time_size = 5 のため、1 回の逆伝播で遡れるのは 5 ステップ分のみです。
これを Truncated BPTT(打ち切り BPTT)と呼びます。
…(以下 stateful表・勾配消失セクションはそのまま残す)
-
アダマール積とも呼びます ↩