seq2seq:系列を読んで、別の系列を書く
入力も出力も、長さのばらばらな系列
翻訳では、「ねこ が さかな を たべる」(5語)を「The cat eats fish」(4語)に変えます。入力も出力も系列で、長さも違います。このような「系列から系列へ」の変換をするモデルを seq2seq(sequence to sequence、系列変換モデル)と呼びます。翻訳のほか、文章の要約、質問への回答、音声から文字への変換などに使われます。
エンコーダとデコーダ
seq2seq は、2つの RNN(ふつうは LSTM)でできています。
- エンコーダ(読む側): 入力の文を1語ずつ読み、最後の隠れ状態に、文全体の内容をまとめる
- デコーダ(書く側): エンコーダの最後の隠れ状態から始めて、出力の文を1語ずつ生成する。生成した単語を次の時刻の入力にして、「文の終わり」の記号を出したら止める
Teacher forcing:学習のときは、正解の単語を入力に使う
デコーダは、本番では自分が生成した単語を次の入力にします。しかし学習の初めは生成がでたらめなので、でたらめな単語を入力にすると、学習がなかなか進みません。そこで学習のときは、自分の生成した単語の代わりに、正解の文の単語 を次の入力に使います。これを Teacher forcing と呼びます。
弱点:すべてを1つのベクトルに詰め込む
エンコーダは、どんなに長い文でも、決まった数(例えば256個)の数の隠れ状態1つに、文全体をまとめなければなりません。短い文ならよいのですが、長い文では、最初の方の内容が薄れて、翻訳の質が落ちてしまいます。前のレッスンの「覚えておく課題」と同じ問題です。
人が翻訳するときは、訳している単語に対応する 元の文の部分を、その都度見返します。「fish」と書くときは「さかな」を、「eats」と書くときは「たべる」を見ます。この「必要なときに、元の文のどこを見るかを決める」仕組みが、次のスライドの Attention です。
Attention:どこに注目するかを、重みで決める
3つの段階
Attention(注意機構)は、デコーダが1語を生成するたびに、エンコーダの すべての時刻の隠れ状態 を見返して、次の3段階で必要な情報を取り出します。使う道具は、これまでの章で学んだものだけです。
- 似ている度合いを測る: デコーダの今の状態(クエリ q、「何を探しているか」)と、エンコーダの各時刻の状態(キー ki、「そこに何があるか」)の 内積(第4章)を計算する
- 割合に変える: 内積を ソフトマックス関数(第7章)に通し、合計が1になる「注目の重み」にする
- 重み付きの平均をとる: 各時刻の情報(バリュー vi)に注目の重みを掛けて足す
注目の重み ai=softmax(q⋅ki),取り出す情報=i∑aivi
数値の例
エンコーダが4つの単語を読み、キーとバリューが次のようになっていたとします(説明のため、キーは2個の数、バリューは1個の数にしています)。デコーダのクエリは q=(2,1) です。
4つの位置への注目の重み。クエリとの内積が一番大きい位置3が0.662で最も大きく、位置1が0.244、位置2が0.090、位置4は0.004とほとんど注目されない。合計は1
クエリと一番似ている位置3(内積3)に、66%の注目が集まりました。取り出す情報は、
0.244×10+0.090×20+0.662×30+0.004×40≈24.3
で、位置3のバリュー(30)に近い値になります。1か所だけを選ぶのではなく、割合で混ぜる ので、ソフトマックスと同じく微分でき、誤差逆伝播で学習できます。
Attention の効果
- デコーダは、1語ごとに、入力のどこでも見返せるので、長い文でも最初の方の情報が薄れにくい
- 注目の重みを見ると、「fish を出すとき、さかな に注目していた」のように、モデルが入力のどこを使ったかが分かる(第9章の説明性にも使われる)
もともと Attention は、seq2seq の RNN を助ける仕組みとして考えられました。ところが、次のレッスンで見るように、Attention だけで 系列を扱う Transformer が生まれ、今の主流になりました。
NumPy で Attention を計算する、よくある誤解と振り返り
import numpy as np
def softmax(z):
e = np.exp(z - z.max()) # 大きな数でも計算が壊れないように、最大値を引いてから
return e / e.sum()
q = np.array([2.0, 1.0]) # クエリ
K = np.array([[1, 0], [0, 1], [1, 1], [-1, 0]], float) # キー(4つの位置)
V = np.array([10.0, 20.0, 30.0, 40.0]) # バリュー
a = softmax(K @ q) # K @ q で、4つのキーとの内積をまとめて計算 → [2, 1, 3, -2]
print(a.round(3)) # → [0.244 0.09 0.662 0.004]
print(a @ V) # 重み付きの平均 → 24.28…
K @ q は、キーの表の各行と q の内積を並べたものです(第6章の @)。a @ V は、重みとバリューの内積で、重み付きの平均になります。
よくある誤解
- 「Attention は、一番似ている1か所だけを選ぶ」: ソフトマックスで割合を決め、すべての位置を混ぜ合わせます(ほとんど0の位置もあります)。だから、微分できて学習できます
- 「注目の重みが大きい場所が、判断の理由のすべて」: 注目の重みは手がかりの1つですが、第9章で見たとおり、1つの見方だけでモデルの判断を説明しきれるわけではありません
- 「Teacher forcing は、本番でも正解を入力に使う」: 本番では正解がないので、自分が生成した単語を次の入力にします。学習のときと本番のときの違いが、誤りの積み重なりの原因になることもあります
振り返り
- seq2seq は、エンコーダで入力の系列を読み、デコーダで出力の系列を1語ずつ生成する。学習では Teacher forcing を使う
- 入力の文を1つのベクトルに詰め込むと、長い文で情報が薄れる
- Attention は、「内積で似ている度合い → ソフトマックスで重み → 重み付きの平均」の3段階で、入力のどこに注目するかを決める。クエリ・キー・バリューの3つを使う