深層学習

系列変換と注意機構

エンコーダが入力系列を固定長ベクトルへ圧縮し、デコーダが出力系列を生成する構成と、そのボトルネックを加法注意で緩和する仕組みを整理する。

  • A|中核
  • 深層学習

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

ひとことで言うと

系列変換(seq2seq)は、入力系列をエンコーダで表現し、デコーダで別の系列へ変換するRNNベースの構成です。基本形は入力全体を1個の固定長ベクトルに押し込みます。加法注意はその圧縮をやめ、デコーダの各ステップで入力側の隠れ状態を重み付きに読み直します。

長い会議録を一枚のメモに要約してから返答を書くのが基本のseq2seqです。注意機構では、返答の各語を書くたびに会議録の該当箇所へ付箋を貼って参照します。メモの容量を増やすのではなく、必要な箇所を都度取り出す点が違います。

なぜ必要か

入力長と出力長が異なる系列を、通常の固定サイズの分類器だけで扱うのは難しいためです。Sutskeverらのseq2seqは、多層LSTMのエンコーダで入力系列を固定次元のベクトルへ写し、別の深いLSTMデコーダで出力系列を復号する方法を示しました。入力と出力の長さを固定する前提を置かず、系列から系列を直接学習する構成です。

ただし、長さの違う入力を同じ1ベクトルへ集約すると、長い入力ほど情報を落としやすくなります。Bahdanauらは、この固定長ベクトルが基本エンコーダ・デコーダの性能向上におけるボトルネックになり得ると考えました。そこで、出力語を予測するたびに入力の関連部分を soft に探索します。エンコーダは入力全体を1個へ完全に要約する負担から解放され、情報を系列上に分散して保持できます。

なお、Sutskeverらは入力文だけを逆順にし、出力文は逆順にしない工夫も報告しています。これにより入力と出力の間に短い依存関係が多く生じ、最適化が容易になった、と論文は説明しています。これは注意機構ではなく、固定長seq2seqの学習を助ける入力順序の工夫です。

固定長ベクトルの問題は「RNNだから必ず長文を扱えない」という意味ではありません。論文が示したのは、入力全体を単一ベクトルに集約する基本構成では、系列が長くなるほど性能上の制約になり得るという点です。

仕組み

入力を x=(x1,…,xTx)\mathbf{x}=(x_1,\ldots,x_{T_x})、エンコーダが各位置で出す隠れ状態(注釈)を hj\mathbf{h}_j、デコーダの直前の隠れ状態を si−1\mathbf{s}_{i-1} とします。hj\mathbf{h}_j は入力の文脈を含み、TxT_x は入力系列長です。デコーダが時刻 ii の出力を作るとき、まず加法的なアライメントモデルで位置 jj の関連度を計算します。

eij=a(si−1,hj)=va⊤tanh⁡(Wamathbfsi−1+Uamathbfhj)e_{ij}=a(\mathbf{s}_{i-1},\mathbf{h}_j)=\mathbf{v}_a^{\top}\tanh\left(W_amathbf{s}_{i-1}+U_amathbf{h}_j\right)

eije_{ij} は出力位置 ii と入力位置 jj の対応度、aa は学習されるフィードフォワードネットワーク、va,Wa,Ua\mathbf{v}_a,W_a,U_a はその重みです。次に、入力位置方向へsoftmaxを適用して確率のような注意重みへ変換します。

αij=exp⁡(eij)∑k=1Txexp⁡(eik)\alpha_{ij}=\frac{\exp(e_{ij})}{\sum_{k=1}^{T_x}\exp(e_{ik})}

αij\alpha_{ij} は時刻 ii に入力位置 jj をどれだけ参照するかを表し、kk はsoftmax内で全入力位置を走査する添字です。最後に、全隠れ状態の重み付き和を文脈ベクトルにします。

ci=∑j=1Txαijhj\mathbf{c}_i=\sum_{j=1}^{T_x}\alpha_{ij}\mathbf{h}_j

ci\mathbf{c}_i は出力位置ごとに再計算される文脈ベクトルです。つまり処理の順序は「スコア eije_{ij} → softmaxで αij\alpha_{ij} → 重み付き和 ci\mathbf{c}_i」です。重みが硬い一箇所の選択ではなく連続値なので、損失からアライメントモデルまで勾配を逆伝播できます。注意機構はこのようにRNNデコーダへ参照経路を追加したもので、ここでは再帰を残しています。後続の構成では注意だけを残して再帰を捨てる方向へ進みますが、それは別のモデルです。

試験でどう問われるか

問われ方正解に寄る条件引っかけ
固定長ベクトルの問題長い入力の情報を1ベクトルへ集約することがボトルネックになるデコーダが入力を一度も参照できない、と一般化する
注意重みの計算順アライメントスコアを計算し、入力位置方向にsoftmaxを適用するsoftmax前のスコアをそのまま重みとする
文脈ベクトルの式全ての hj\mathbf{h}_j を αij\alpha_{ij} で重み付けして足す最大スコアの状態だけを選ぶhard alignmentと混同する
インデックスの意味ii は出力位置、jj は入力位置。重みは出力ステップごとに変わる系列全体で1組の重みを使い回す
Sutskeverらの工夫入力だけを逆順にし、短期依存を増やして最適化を容易にした入力と出力の両方を逆順にしたとする

実装で確かめる

小さな配列で、スコア、softmax、文脈ベクトルの形をそのまま実行します。ここでは学習済みの重みではなく、加法注意の計算経路だけを確認します。

import numpy as np

h = np.array([[1., 0.], [0., 2.], [1., 1.]])  # 入力3位置の隠れ状態 h_j
s = np.array([.5, 1.])                       # デコーダの直前状態 s_{i-1}
Wa = np.eye(2); Ua = np.eye(2); va = np.array([1., -1.])
e = np.array([va @ np.tanh(Wa @ s + Ua @ hj) for hj in h])
alpha = np.exp(e - e.max()); alpha /= alpha.sum()
c = alpha @ h
print("alpha sum:", alpha.sum())
print("context:", np.round(c, 6))

alpha の総和は1になり、context は3個の隠れ状態を同じ次元のまま混ぜたベクトルになります。実装で次元が合わないときは、スコアが入力位置ごとに1個、重みが入力系列方向に正規化され、文脈が ∑jαijhj\sum_j\alpha_{ij}\mathbf{h}_j になっているかを確認します。

取り違えやすいもの

用語系列変換・加法注意との切り分け
固定長seq2seqエンコーダの最後の表現をデコーダへ渡し、入力全体を1ベクトルに圧縮する基本形です
加法注意デコーダ状態と各エンコーダ状態からスコアを作り、全状態の重み付き和を各ステップで作ります
hard alignment入力位置を離散的に1つ選ぶ考え方です。加法注意のsoftな重みとは異なります
Self-Attention系列内の要素同士を参照する注意の枠組みです。ここで扱う加法注意はRNNのエンコーダ・デコーダに組み込まれています
Transformer再帰を使わず注意を中心に構成する別モデルです。「注意を使う」だけで加法注意と同一視しません

想起チェック

基本のseq2seqで入力系列が情報ボトルネックになる場所はどこか

エンコーダが入力全体を1個の固定長ベクトルへ押し込む場所です。入力が長いほど、同じ容量へ多くの情報を集約する必要があります。

加法注意の計算を3段階で答えると

各入力位置のアライメントスコア eije_{ij} を計算し、softmaxで αij\alpha_{ij} に変換し、全隠れ状態の重み付き和 ci=∑jαijhj\mathbf{c}_i=\sum_j\alpha_{ij}\mathbf{h}_j を作ります。

添字 i と j はそれぞれ何の位置を表すか

ii はデコーダが生成する出力位置、jj はエンコーダが読んだ入力位置です。したがって、出力位置 ii ごとに入力方向の注意重みが変わります。

Sutskeverらが報告した入力逆順の工夫は、何を逆順にしたか

入力系列だけを逆順にし、出力系列は逆順にしていません。入力と出力の短期依存を増やし、最適化を容易にする意図です。

出典