深層学習

BPTT

RNNを時間方向に展開し、各時刻の損失から共有パラメータへ逆向きに勾配を集める手続きと、長い系列を区切るtruncated BPTTを整理します。

  • B|標準
  • 深層学習

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

ひとことで言うと

BPTT(Backpropagation Through Time)は、RNNを時刻ごとに横へ展開した計算グラフに、通常の誤差逆伝播を適用する手続きです。各時刻で同じ重みを使うため、逆向きにたどりながら得た勾配を時刻方向に足し合わせてから更新します。

同じ担当者が毎日の判断を行う業務記録を、日付ごとの別担当者として並べて検証するようなものです。最後の判断の誤りを前の日へ戻し、各日の担当者が共通の手順にどれだけ関与したかを合算します。

なぜ必要か

通常の逆伝播は層の順序を逆にたどりますが、RNNの依存関係には時刻もあります。時刻 tt の状態は t−1t-1 の状態を使うため、系列の最後で生じた損失の原因を前の時刻まで帰属させるには、時間方向へグラフを展開する必要があります。

展開長を TT とすると、順伝播も逆伝播も基本的に TT 回分のセル計算を行います。さらに各時刻の状態や活性化を逆伝播用に保持するため、計算量とメモリ使用量は展開長に比例して増えます。長い系列をそのまま一括処理できないとき、精度だけでなくこの二つの資源が区切り方を決めます。

展開方法逆伝播の範囲主な負担
完全なBPTT系列全体長いほど計算量と保存状態が増える
truncated BPTT固定長の窓だけ境界より前へ勾配が届かない

仕組み

入力を xt\mathbf{x}_t、隠れ状態を ht\mathbf{h}_t、共有される再帰重みを WhhW_{hh}、入力重みを WxhW_{xh}、バイアスを b\mathbf{b}、状態更新の活性化関数を ff とします。時刻 tt の状態と損失 LtL_t は次のように計算します。

ht=f(Whhht−1+Wxhxt+b),L=∑t=1TLt\mathbf{h}_t=f(W_{hh}\mathbf{h}_{t-1}+W_{xh}\mathbf{x}_t+\mathbf{b}),\qquad L=\sum_{t=1}^{T}L_t

TT は展開する時刻数、LtL_t は時刻 tt の損失です。逆伝播では LtL_t の影響を ht−1\mathbf{h}_{t-1} へ戻し、同じ WhhW_{hh} が各時刻で使われた分の寄与を合計します。概念的には、状態に対する勾配を δt=∂L/∂ht\boldsymbol{\delta}_t=\partial L/\partial\mathbf{h}_t とすると、次の時刻から来る項を含めて計算します。

δt=∂Lt∂ht+WhhTδt+1⊙f′(at),∂L∂Whh=∑t=1Tδtht−1T\boldsymbol{\delta}_t=\frac{\partial L_t}{\partial\mathbf{h}_t}+W_{hh}^{\mathsf{T}}\boldsymbol{\delta}_{t+1}\odot f'(\mathbf{a}_t),\qquad \frac{\partial L}{\partial W_{hh}}=\sum_{t=1}^{T}\boldsymbol{\delta}_t\mathbf{h}_{t-1}^{\mathsf{T}}

at\mathbf{a}_t は活性化前の値、f′f' はその導関数、⊙\odot は要素ごとの積です。重要なのは、時間ごとに別の重みを学習するのではなく、共有重みへの勾配を全時刻から集める点です。

truncated BPTTでは、系列を長さ KK の窓に分け、各窓の中だけを逆伝播します。窓の先頭で勾配を止めるため、前窓から渡した状態を数値として使えても、その状態を作った計算グラフには戻りません。したがって勾配は境界を越えて伝わらず、KK より前の入力が現在の損失へ与える勾配はその更新では計算されません。

試験でどう問われるか

問われ方正解に寄る条件引っかけ
BPTTの対象RNNを時間方向に展開し、通常の逆伝播を適用する時刻ごとに別モデルを学習すると考える
計算量・メモリ展開長に比例して増える。状態保存も必要共有重みだから時刻数に依存しないとする
truncated BPTT長さ KK の範囲で逆伝播し、境界で勾配を切る状態を渡せば勾配も前窓へ届くとする
状態の扱い窓間で状態を引き継ぐか、初期化するかを目的に合わせて選ぶ状態を引き継ぐことと計算グラフをつなぐことを同一視する

実装で確かめる

窓ごとに状態を渡す実装と、窓ごとに状態をゼロへ戻す実装は、順方向の状態の扱いが異なります。自動微分ライブラリでは、引き継ぐ状態を detach してから次の窓へ渡すと、値だけを継承してBPTTの境界を作れます。

import numpy as np

def run_chunks(xs, W, U, chunk, carry_state):
    h = np.zeros(W.shape[0])
    outputs = []
    for start in range(0, len(xs), chunk):
        if not carry_state:
            h = np.zeros_like(h)
        for x in xs[start:start + chunk]:
            h = np.tanh(W @ h + U @ x)
            outputs.append(h.copy())
        # NumPyではグラフを持たないが、ここがdetachする境界に相当する
    return np.asarray(outputs)

rng = np.random.default_rng(0)
xs = rng.normal(size=(6, 2))
W, U = rng.normal(size=(3, 3)), rng.normal(size=(3, 2))
assert run_chunks(xs, W, U, 2, True).shape == (6, 3)

取り違えやすいもの

用語BPTTとの切り分け
通常の逆伝播層方向の計算グラフを逆にたどる。BPTTはその考え方を時間方向へ適用する
truncated BPTTBPTTの近似的な実行方法。指定窓の外へ勾配を流さない
状態リセット次の窓の初期状態をゼロなどにする実装。勾配を切ることとは別の選択
勾配クリッピング勾配の大きさを制限する処理。逆伝播をどこまで行うかは決めない

想起チェック

BPTTで何を時間方向に行うか

RNNの計算グラフを時刻ごとに展開し、そのグラフへ逆伝播を適用します。

展開長を増やすと何が増えるか

セル計算の回数と、逆伝播のために保持する状態の量が増えます。

状態を引き継いだtruncated BPTTで、勾配も境界を越えるか

越えません。状態の数値は次の窓へ渡せますが、境界より前の計算グラフを切れば、その範囲へは勾配が伝わりません。

出典