Hacnosuke
Title 12. AttentionとTransformer 機械学習 /

AttentionとTransformer

自然言語をはじめとした系列データを処理するニューラルネットワークとして、RNNやLSTMが広く使われてきた。しかし、これらのモデルでは情報を系列に沿って順番に伝える必要がある。例えば、 という系列において、 の情報を の処理に活かすためには、中間の隠れ状態を通して情報を伝播させなければならない。そのため、要素同士が遠く離れるほど、その関連性を学習することが難しくなる。また、系列を順番に処理する必要があるため、計算を並列化しにくいという問題もある。

これに対してAttentionでは、情報を隠れ状態によって順番に伝えるのではなく、系列中の各要素を直接参照し、その中から現在の処理に重要な情報を選択する。これにより、系列中で離れた要素の情報も、途中の隠れ状態を何度も経由せずに利用できる。つまり、 の処理において、 の情報を直接参照し、重みをつけて集約することができる。これにより、長期的な依存関係を学習しやすくなる。

Transformerは、このAttentionを系列処理の中心に据え、RNNのような再帰構造を取り除いたモデルである。系列全体の関係をまとめて計算することで、長期的な依存関係を扱いやすくすると同時に、学習時の並列計算を可能にした。以降では、Transformerで使われるAttentionの仕組みを具体的に説明する。

Attention

まず、系列データを以下のように表す。

は系列の各要素を表すベクトルであり、 はモデルの隠れ状態の次元数である。ここでの目標は、 番目の要素 に対応する出力 を次のようにすべての系列要素から集約して計算することである。

は系列の各要素を変換する関数であり、 は が をどれだけ参照するかを表す重みである。すなわち、 が大きいほど、 の情報が に強く反映される。Attentionでは をバリュー、 をアテンション重みと呼び、これはクエリとキーの適合度にもとづいて計算される。以下、以下のように定義する。

  • クエリ:
  • キー:
  • バリュー:

これを元に書き下してみると、

となる。ここで はクエリとキーの類似度を表す関数である。さて、一般的に"sim"は内積で定義されることが多い。すなわち、 である。これにより、クエリとキーの類似度が大きいほど、バリューの重みが大きくなる。また、以下に二つの工夫を追加する。

  • 重み は正規化される必要がある。つまり、 である必要がある。そこで、ソフトマックス関数を用いて正規化する。すなわち、
  • クエリとキーの内積は、ベクトルの次元数が大きくなると値が大きくなりやすい。これにより、ソフトマックス関数を通した際に勾配が小さくなり、学習が遅くなることがある。そこで、内積を で割ることで、値のスケールを調整する。

以上より、Attentionの出力は以下のように計算される。

これは簡単な行列計算で表すことができる。まず、クエリ、キー、バリューをそれぞれ行列に集約する。

すると、正規化する前の類似度行列は

となる。行列に対するソフトマックス関数は、一般的に各行ごとに正規化を行うことで決められているので、Attentionの出力は以下のようにまとめられる。

Scaled Dot-Product Attentionの計算手順(Vaswani et al., 2017 Figure 2左)

ここでは、クエリ、キー、バリューは一つの同じ系列から生成されるSelf-Attentionを説明してきた。この場合、これらは次のように計算される。

ここではクエリとキーは対称的な関係を持っている。しかし、実際にはそうではない。クエリを別の系列データから生成するようにしてみよう。

すると は、クエリ側の系列 と、キー側の系列 の間の類似度を表す行列となる。これにより、1つの要素の出力は

となる。ここではどちらの系列から来たのかを明示するために、クエリ側の系列には上付きの 、キー・バリュー側の系列には上付きの を付けている。この式は出力に対応する位置のクエリをキーとの関連度を元にバリューで重み付けして集約することを意味している。いわば、クエリ系列をキー・バリュー系列に照らし合わせて処理することになる。これをCross-Attentionと呼ぶ。

Transformer

Transformerは、RNNやCNNに頼らず、Attention機構を中心に系列を処理するエンコーダ・デコーダ型のモデルである。各ブロックは、Multi-Head Attention、位置ごとのFeed-Forward Network、Residual Connection、Layer Normalizationから構成される。 エンコーダは入力列から文脈表現を生成し、デコーダはその文脈と既に生成した出力を用いて次のトークンを予測する。

Transformerのエンコーダ・デコーダ構造(Vaswani et al., 2017 Figure 1)

Self-AttentionとCross-Attention

Transformerのエンコーダとデコーダのブロックには、Self-AttentionとCross-Attentionの二種類のAttentionが存在する。Self-AttentionはAttentionの入力であったクエリ、キー、バリューが全て同じ入力 から生成される。

Self-Attentionは、入力 の各トークン位置どうしの関係を見ながら表現を更新するためのもので、エンコーダとデコーダの両方で用いられる。 一方で、Cross-Attentionはクエリが入力 とは異なる入力 から生成される。

Cross-Attentionは、クエリを生成した入力 を用いて、 の情報を集約するためのもので、デコーダがエンコーダ出力を参照しながら次のトークンを生成するために使われる。

マルチヘッドAttention

TransformerのAttentionは、単一のAttentionではなく、複数のAttentionを並列に配置したマルチヘッドAttention(Multi-Head Attention)である。これはクエリ、キー、バリューのベクトルを複数の小さなベクトルに線形射影して、それぞれに独立したAttentionを適用し、最後にそれらの出力を結合するものである。マルチヘッドAttentionは、Transformerのエンコーダとデコーダの両方で使われる。

マスク付きAttention

学習時のデコーダは正解系列をまとめて受け取れるが、そのままでは未来のトークンを見てしまい、推論時との条件がずれてしまう。そのため、未来位置に対応するスコアに を加え、Softmax後の重みが になるようにする。具体的には、以下のような行列を用いて

出力を以下のように計算する。

すると、重み行列のうち未来位置に対応する部分は となり、未来のトークンを参照できなくなる。Masked Attentionは、Transformerのデコーダ内のSelf-Attentionで使用される。

位置エンコーディング

Attentionは、RNNやLSTMのように順序を逐次的には処理しないため、そのままではトークンの並び順を区別できない。そこでTransformerは、各トークン表現にその位置を表すベクトルを加える。原論文では次のような正弦波型の位置エンコーディングが使われた。

ここで、 は系列内の位置を表す整数であり、 や はベクトルのインデックスである。 はトークン表現の次元数である。位置エンコーディングは入力埋め込み に加えて使用される。

Position Encodingは以下の点で優れている。

  • はトークン内容 に依存せず、位置だけから定まる。
  • 正弦・余弦の加法定理により、固定オフセット に対する を の線形関数として表せる。
  • Vaswaniらは、この設計が学習時より長い系列への外挿を助ける可能性を期待して、学習型ではなく正弦波型の位置エンコーディングを採用した。

FFN(Feed-Forward Network)

Attentionは位置間の情報を混合するが、各位置に対する非線形変換は別途必要である。そのためTransformerは、各位置に同じ2層の全結合ネットワークを独立に適用するフィードフォワードネットワーク(Feed-Forward Network; FFN)を配置している。原論文ではReLUを用いた次の形で定義される。

FFNは他の位置の情報を参照せず、各トークン位置の表現を個別に変換する。Transformerではこの層もエンコーダとデコーダの両方に含まれる。