本文へスキップ
ひもとくAI

RNN:前の情報を隠れ状態にまとめて渡していく

レッスン 2/7

RNN:前の情報を「隠れ状態」にまとめて渡していく

本を読むとき、前の内容を覚えながら読む

文章を1語ずつ読むとき、私たちは、それまでに読んだ内容を頭の中に覚えておき、新しい単語を読むたびに、その記憶を少しずつ書き換えています。RNN(リカレントニューラルネットワーク、再帰型ニューラルネットワーク)は、この読み方をまねたモデルです。

RNN は、それまでに読んだ内容を、隠れ状態 hh という数のベクトル(この例では16個の数)にまとめておきます。そして、系列の1つ1つの位置(時刻 tt)で、次の計算をします。

ht=tanh⁡(xtWx+ht−1Wh+b)h_t = \tanh(x_t W_x + h_{t-1} W_h + b)
  • xtx_t: 時刻 tt の入力(単語の埋め込みベクトルや、その時刻の値)
  • ht−1h_{t-1}: 1つ前の時刻の隠れ状態(最初は全部0)
  • WxW_x、WhW_h、bb: 重みとバイアス
  • tanh⁡\tanh: 第6章で紹介した活性化関数(−1〜1の値にする)

第6章の「入力 @ 重み + バイアス → 活性化関数」に、「1つ前の隠れ状態 @ 重み」を足しただけ です。新しい隠れ状態 hth_t は、次の時刻の計算に渡されます。

同じ重みを、どの時刻でも使う

大切なのは、WxW_x、WhW_h、bb が、どの時刻でも同じ であることです。

  • 系列の長さがいくつでも、同じ計算を繰り返すだけで処理できる(長さがばらばらでもよい)
  • 「最初の位置で学んだこと」を、ほかの位置でも使える

系列を最後まで読んだら、最後の隠れ状態 hTh_T から、第6章と同じように出力(分類の確率や、予測の値)を計算します。

時間方向に広げて見る

RNN の計算は、同じ層を時刻の数だけ横に並べた、長いネットワークだと見ることもできます。

入力時刻1時刻2時刻3出力
RNN を時間方向に広げた図(説明のための模式図)。時刻1・2・3の隠れ状態(それぞれ4個の数)が左から右に並び、各時刻の隠れ状態は、次の時刻の隠れ状態の計算に使われる。実際には、各時刻に入力も加わり、どの時刻でも同じ重みを使う。最後の隠れ状態から出力を計算する

学習(誤差逆伝播)も、この広げたネットワークで、出力の側から時刻をさかのぼって勾配を計算します。これを BPTT(Backpropagation Through Time、時間をさかのぼる誤差逆伝播)と呼びます。同じ重みがどの時刻でも使われているので、各時刻で計算した勾配を 足し合わせて、重みを更新します。

実験:NumPy で RNN を作り、波の続きを予測する

NumPy で書いた RNN の順伝播

import numpy as np

def rnn_forward(X, Wx, Wh, b):
    # X: (データの数, 時刻の数, 入力の数)
    h = np.zeros((X.shape[0], Wh.shape[0]))       # 最初の隠れ状態は0
    for t in range(X.shape[1]):                   # 時刻の順に
        h = np.tanh(X[:, t] @ Wx + h @ Wh + b)    # 前の h と、今の入力から、新しい h
    return h                                      # 最後の隠れ状態

X[:, t] は、すべてのデータの時刻 tt の入力です。第6章と同じく、たくさんのデータをまとめて計算しています。この関数は、PyTorch の nn.RNN と同じ値を出すことを確かめています。

実験:サイン波の次の値を当てる

小さな雑音を加えたサイン波(なめらかな波)で、「直前の10個の値から、次の値を当てる」課題を作りました。隠れ状態16個の RNN を、BPTT(自分で書いた勾配の計算。数値微分と一致することを確かめた)と勾配降下法で300回学習させました(ブラウザと同じ環境で1.3秒)。

時刻(テストデータ)値01020304050-101
  • 本当の値
  • RNN の予測
テストデータの最初の60時刻の、本当の値(実線)と RNN の予測(破線)。どちらも−1から1の間をなめらかに上下する波で、予測の波は本当の波にほぼ重なっている。本当の値は雑音で細かくぎざぎざしているが、予測はなめらか
予測の方法テストデータの二乗誤差
何も学習しない(いつも0と答える)0.486
直前の値をそのまま答える0.026
RNN0.006
直前の10個に重みを掛けて足す(第4章の線形回帰)0.004

RNN は、波の形をよく学習しました。ところが、ずっと単純な 線形回帰の方が、少しよい結果 でした。なめらかな波は、直前のいくつかの値の足し算でよく表せるので、この課題には複雑なモデルは必要なかったのです。

新しいモデルを試すときは、いつも 単純な方法と比べる(ベースラインを置く)ことが大切です。RNN が本当に役に立つのは、単純な足し算では表せない、もっと複雑な系列(文章など)です。

PyTorch の RNN、双方向RNN、振り返り

PyTorch で RNN を使う

import torch
from torch import nn

rnn = nn.RNN(input_size=1, hidden_size=16, batch_first=True)   # 入力1個、隠れ状態16個
out = nn.Linear(16, 1)                                           # 最後の隠れ状態 → 予測

X = torch.randn(32, 10, 1)          # (データの数, 時刻の数, 入力の数)
hs, h_last = rnn(X)                 # hs: すべての時刻の隠れ状態、h_last: 最後の隠れ状態
pred = out(hs[:, -1])               # 最後の時刻の隠れ状態から予測
  • batch_first=True で、NumPy の例と同じ「データの数, 時刻の数, 入力の数」の並びになります
  • 勾配の計算(BPTT)は、PyTorch が自動で行います(第7章の自動微分)
  • PyTorch の式は tanh⁡(xWih⊤+bih+hWhh⊤+bhh)\tanh(x W_{ih}^{\top} + b_{ih} + h W_{hh}^{\top} + b_{hh}) で、バイアスが2つある点だけが NumPy の例と違います(足せば同じ)

双方向RNN:後ろからも読む

文章の途中の単語の意味は、前だけでなく、後ろの単語で決まることもあります(「はし を わたる」の「はし」が橋か箸かは、後ろを読むと分かる)。前から読む RNN と、後ろから読む RNN を両方使い、2つの隠れ状態を並べて使うモデルを 双方向RNN と呼びます(nn.RNN(..., bidirectional=True))。ただし、文章を先頭から1語ずつ生成するような、「まだ後ろがない」場面では使えません。

よくある誤解

  • 「RNN は、時刻ごとに別々の重みを持っている」: どの時刻でも同じ重みを使います。だから、長さの違う系列を同じモデルで扱えます
  • 「系列データには、いつも RNN が一番よい」: サイン波の実験では、線形回帰の方がよい結果でした。単純な方法と比べてから選びます
  • 「隠れ状態に、それまでの入力がすべてそのまま記録されている」: 隠れ状態は、決まった数(例では16個)の数に、それまでの情報を まとめたもの です。長い系列では、前の情報ほど薄れていきます(次のレッスン)

振り返り

  • RNN は、隠れ状態 ht=tanh⁡(xtWx+ht−1Wh+b)h_t = \tanh(x_t W_x + h_{t-1} W_h + b) を、時刻の順に計算して、前の情報を次へ渡していく
  • どの時刻でも同じ重みを使う。学習は、時刻をさかのぼる誤差逆伝播(BPTT)で、各時刻の勾配を足し合わせる
  • 双方向RNNは、前からと後ろからの両方で読む
  • 新しいモデルは、単純な方法(ベースライン)と比べて評価する

演習

演習を読み込んでいます…