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

seq2seq と Attention:どこに注目するかを重みで決める

レッスン 4/7

seq2seq:系列を読んで、別の系列を書く

入力も出力も、長さのばらばらな系列

翻訳では、「ねこ が さかな を たべる」(5語)を「The cat eats fish」(4語)に変えます。入力も出力も系列で、長さも違います。このような「系列から系列へ」の変換をするモデルを seq2seq(sequence to sequence、系列変換モデル)と呼びます。翻訳のほか、文章の要約、質問への回答、音声から文字への変換などに使われます。

エンコーダとデコーダ

seq2seq は、2つの RNN(ふつうは LSTM)でできています。

  1. エンコーダ(読む側): 入力の文を1語ずつ読み、最後の隠れ状態に、文全体の内容をまとめる
  2. デコーダ(書く側): エンコーダの最後の隠れ状態から始めて、出力の文を1語ずつ生成する。生成した単語を次の時刻の入力にして、「文の終わり」の記号を出したら止める

Teacher forcing:学習のときは、正解の単語を入力に使う

デコーダは、本番では自分が生成した単語を次の入力にします。しかし学習の初めは生成がでたらめなので、でたらめな単語を入力にすると、学習がなかなか進みません。そこで学習のときは、自分の生成した単語の代わりに、正解の文の単語 を次の入力に使います。これを Teacher forcing と呼びます。

弱点:すべてを1つのベクトルに詰め込む

エンコーダは、どんなに長い文でも、決まった数(例えば256個)の数の隠れ状態1つに、文全体をまとめなければなりません。短い文ならよいのですが、長い文では、最初の方の内容が薄れて、翻訳の質が落ちてしまいます。前のレッスンの「覚えておく課題」と同じ問題です。

人が翻訳するときは、訳している単語に対応する 元の文の部分を、その都度見返します。「fish」と書くときは「さかな」を、「eats」と書くときは「たべる」を見ます。この「必要なときに、元の文のどこを見るかを決める」仕組みが、次のスライドの Attention です。

Attention:どこに注目するかを、重みで決める

3つの段階

Attention(注意機構)は、デコーダが1語を生成するたびに、エンコーダの すべての時刻の隠れ状態 を見返して、次の3段階で必要な情報を取り出します。使う道具は、これまでの章で学んだものだけです。

  1. 似ている度合いを測る: デコーダの今の状態(クエリ qq、「何を探しているか」)と、エンコーダの各時刻の状態(キー kik_i、「そこに何があるか」)の 内積(第4章)を計算する
  2. 割合に変える: 内積を ソフトマックス関数(第7章)に通し、合計が1になる「注目の重み」にする
  3. 重み付きの平均をとる: 各時刻の情報(バリュー viv_i)に注目の重みを掛けて足す
注目の重み ai=softmax(q⋅ki),取り出す情報=∑iaivi\text{注目の重み } a_i = \text{softmax}(q \cdot k_i), \qquad \text{取り出す情報} = \sum_{i} a_i v_i

数値の例

エンコーダが4つの単語を読み、キーとバリューが次のようになっていたとします(説明のため、キーは2個の数、バリューは1個の数にしています)。デコーダのクエリは q=(2,1)q = (2, 1) です。

単語の位置1234
キー kik_i(1, 0)(0, 1)(1, 1)(−1, 0)
内積 q⋅kiq \cdot k_i213−2
注目の重み(ソフトマックス)0.2440.0900.6620.004
バリュー viv_i10203040
単語の位置0.24410.0920.66230.0044
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.30.244 \times 10 + 0.090 \times 20 + 0.662 \times 30 + 0.004 \times 40 \approx 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 は、キーの表の各行と qq の内積を並べたものです(第6章の @)。a @ V は、重みとバリューの内積で、重み付きの平均になります。

よくある誤解

  • 「Attention は、一番似ている1か所だけを選ぶ」: ソフトマックスで割合を決め、すべての位置を混ぜ合わせます(ほとんど0の位置もあります)。だから、微分できて学習できます
  • 「注目の重みが大きい場所が、判断の理由のすべて」: 注目の重みは手がかりの1つですが、第9章で見たとおり、1つの見方だけでモデルの判断を説明しきれるわけではありません
  • 「Teacher forcing は、本番でも正解を入力に使う」: 本番では正解がないので、自分が生成した単語を次の入力にします。学習のときと本番のときの違いが、誤りの積み重なりの原因になることもあります

振り返り

  • seq2seq は、エンコーダで入力の系列を読み、デコーダで出力の系列を1語ずつ生成する。学習では Teacher forcing を使う
  • 入力の文を1つのベクトルに詰め込むと、長い文で情報が薄れる
  • Attention は、「内積で似ている度合い → ソフトマックスで重み → 重み付きの平均」の3段階で、入力のどこに注目するかを決める。クエリ・キー・バリューの3つを使う

演習

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