深層学習

GRU

GRUは更新ゲートとリセットゲートで隠れ状態を直接制御する、セル状態を持たないゲート付きRNNです。

  • B|標準
  • 深層学習

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

ひとことで言うと

GRU(Gated Recurrent Unit)は、LSTMに動機づけられたゲート付きRNNです。更新ゲートとリセットゲートの2つで過去の隠れ状態を残す量と、候補状態を作るときに過去を参照する量を調整します。LSTMのような独立したセル状態は持たず、隠れ状態だけを引き継ぎます。

一つのメモ帳だけを使い、更新ゲートで「前のページをどれだけ残すか」、リセットゲートで「新しいページを書くとき前のページをどれだけ読むか」を決める方式です。保管場所を隠れ状態に一本化するので、状態の対応を追いやすくなります。

なぜ必要か

通常のRNNでは、各時刻の隠れ状態を非線形変換で作り直すため、長い系列の情報を学習で保ちにくいという問題があります。GRUは、過去をそのまま混ぜて次へ渡す経路をゲートで作り、長期情報を残す判断をデータから学習できるようにしました。

原論文は、提案した隠れユニットをLSTMより計算・実装しやすいものとして説明しています。ゲートが2つで、セル状態も別に持たないため、同じ隠れサイズなら管理する状態とゲート由来の計算が少なく、一般にパラメータ数と計算量を抑えやすいのが実装上の利点です。ただし実測速度はカーネルやバッチ形状にも左右されます。GRUとLSTMのどちらが常に優れるかは、タスクや設定を離れて決着した話ではありません。

不便GRUで導入するもの得られる制御
過去を毎回作り直す更新ゲート隠れ状態を保持する量
古い情報を候補に混ぜ続けるリセットゲート過去を参照する量

仕組み

入力を xt\mathbf{x}_t、前時刻の隠れ状態を ht−1\mathbf{h}_{t-1}、シグモイド関数を σ\sigma、要素積を ⊙\odot とします。まずリセットゲート rt\mathbf{r}_t と更新ゲート zt\mathbf{z}_t を計算します。

rt=σ(Wrxt+Urht−1),zt=σ(Wzxt+Uzht−1)\mathbf{r}_t=\sigma(\mathbf{W}_r\mathbf{x}_t+\mathbf{U}_r\mathbf{h}_{t-1}),\qquad \mathbf{z}_t=\sigma(\mathbf{W}_z\mathbf{x}_t+\mathbf{U}_z\mathbf{h}_{t-1})

リセットゲートは候補状態を作るときの過去の参照量です。候補状態 h~t\tilde{\mathbf{h}}_t は、過去の隠れ状態に rt\mathbf{r}_t を掛けてから変換します。ϕ\phi は通常 tanh⁡\tanh を使う非線形関数です。

h~t=ϕ(Wxt+U(rt⊙ht−1))\tilde{\mathbf{h}}_t=\phi\left(\mathbf{W}\mathbf{x}_t+\mathbf{U}(\mathbf{r}_t\odot\mathbf{h}_{t-1})\right)

最後に更新ゲートで、過去を残す側と候補を採用する側を混ぜます。

ht=zt⊙ht−1+(1−zt)⊙h~t\mathbf{h}_t=\mathbf{z}_t\odot\mathbf{h}_{t-1}+(1-\mathbf{z}_t)\odot\tilde{\mathbf{h}}_t

zt\mathbf{z}_t が1に近ければ過去を多く保持し、0に近ければ候補へ更新します。rt\mathbf{r}_t が0に近い部分では、候補状態が過去をほぼ無視します。試験では更新式の係数の向き、リセットを掛ける位置、セル状態を別記しない点が確認箇所です。

試験でどう問われるか

問われ方正解に寄る条件引っかけ
ゲートの役割更新は過去を残す量、リセットは候補計算で過去を読む量2つの役割を逆にする
状態の構造隠れ状態を更新し、独立したセル状態を持たないLSTMと同じ2状態とする
更新式の読み取りzt\mathbf{z}_t と 1−zt1-\mathbf{z}_t が過去と候補を分担する両方に同じ係数を掛ける
LSTMとの比較構造が簡素で、優劣はタスク依存GRUが常に高精度・高速と断定する

実装で確かめる

次のコードは、1時刻ぶんのGRU更新をNumPyでそのまま計算します。z が大きいと、出力が前の隠れ状態に近づくことを確認できます。

import numpy as np

def gru_step(x, h, Wz, Uz, Wr, Ur, W, U):
    sigmoid = lambda a: 1 / (1 + np.exp(-a))
    z = sigmoid(Wz @ x + Uz @ h)
    r = sigmoid(Wr @ x + Ur @ h)
    h_tilde = np.tanh(W @ x + U @ (r * h))
    return z * h + (1 - z) * h_tilde, z, r

rng = np.random.default_rng(0)
x, h = rng.normal(size=3), rng.normal(size=4)
M = [rng.normal(size=(4, 3)), rng.normal(size=(4, 4))]
new_h, z, r = gru_step(x, h, M[0], M[1], M[0], M[1], M[0], M[1])
print(new_h.shape, z.min() >= 0 and z.max() <= 1, r.min() >= 0 and r.max() <= 1)

取り違えやすいもの

用語GRUとの切り分け
LSTMゲートとセル状態を持つ。GRUは隠れ状態に状態を統合する
通常のRNNゲートなしで候補をほぼ毎回更新する。GRUは保持・リセットを学習する
更新ゲート過去の隠れ状態と候補の混合比を決める
リセットゲート候補状態の計算で過去の隠れ状態をどれだけ使うか決める

想起チェック

GRUがLSTMと異なり別のセル状態を持たない理由は

GRUは隠れ状態だけを更新し、更新ゲートで保持経路を作ります。状態を隠れ状態に統合した構造です。

更新ゲートが1に近いと、隠れ状態はどうなるか

ht=zt⊙ht−1+(1−zt)⊙h~t\mathbf{h}_t=\mathbf{z}_t\odot\mathbf{h}_{t-1}+(1-\mathbf{z}_t)\odot\tilde{\mathbf{h}}_t なので、前時刻の隠れ状態を多く保持します。

リセットゲートが0に近いと候補状態の計算はどう変わるか

rt⊙ht−1\mathbf{r}_t\odot\mathbf{h}_{t-1} が小さくなり、候補状態は現在入力を中心に、過去をほぼ無視して計算されます。

出典