マルチヘッドアテンション
導入1
それぞれクエリ・キー・バリューと呼ばれる三つの行列$\mathbf{Q} \in M^{d_{k} \times n}$、$\mathbf{K} \in M^{d_{k} \times m}$、$\mathbf{V} \in M^{d_{v} \times m}$と、クエリとキーに対するスコア関数$f : M^{d_{k} \times n} \times M^{d_{k} \times m} \to M^{m \times n}$が与えられたとしよう。アテンションとは、次のように定義される関数をいう。
$$ \begin{align*} \operatorname{Attention}: M^{d_{k} \times n} \times M^{d_{k} \times m} \times M^{d_{v} \times m} &\to M^{d_{v} \times n} \\ (\mathbf{Q},\mathbf{K},\mathbf{V}) &\mapsto \mathbf{V} \operatorname{Softmax}\left(f(\mathbf{Q},\mathbf{K})\right) \tag{1} \end{align*} $$
ここで$\operatorname{Softmax}$とは、与えられた行列$\mathbf{X} = \begin{bmatrix} \mathbf{x}_{1} & \cdots & \mathbf{x}_{N}\end{bmatrix}$に対して、各列ベクトルにソフトマックス $\operatorname{softmax}$を取り、各列の和が$1$になるようにする関数をいう。
$$ \operatorname{Softmax}(\mathbf{X}) := \begin{bmatrix} \underset{\vert}{\overset{\vert}{\operatorname{softmax}(\mathbf{x}_{1})}} & \cdots & \underset{\vert}{\overset{\vert}{\operatorname{softmax}(\mathbf{x}_{N})}} \end{bmatrix} $$
このように計算されたアテンションの関数値は、クエリとキーの適合度を重みとしてバリューを加重平均した行列である(詳しい説明はアテンションの文書を参照せよ)。
$$ \operatorname{Attention} (\mathbf{Q},\mathbf{K},\mathbf{V}) \in M^{d_{v} \times n} $$
ここで少しの間、画像データと畳み込みニューラルネットワーク(CNN)について考えてみよう。サイズが$(1, N, N)$の入力画像は、CNNを通してサイズが$(C, M, M)$の配列へと加工される。このとき$C$をチャンネルの次元と呼び、各チャンネルには入力データに対する抽象的な情報が込められている。この抽象的な情報が色であるならば、カラー写真はRGBの三つのチャンネルを持ち$(3, N, N)$のように表現され、白黒写真はチャンネルが一つだけなので$(1, N, N)$のように表現される。もちろん、CNNを通過して得られる$(C, M, M)$配列のチャンネルに込められた情報は、このように人間が直観的に分かるものではないが、ともかく入力データの情報が$C$種類に分解されたと理解することができる。同様に、アテンションを複数使って$\mathbf{Q}$、$\mathbf{K}$、$\mathbf{V}$の情報を複数のチャンネルで表現することができるだろう。アテンションでは、このようなチャンネルをヘッドheadと呼ぶ。
ところが$(1)$を見ると、$\operatorname{Attention}$関数自体は学習可能なパラメータを含まない。ただソフトマックスと行列積のみで構成されている(もちろん$f$をニューラルネットワークで構成すれば学習可能なパラメータが含まれるが、ここではそうではないとしよう)。そこで、クエリ、キー、バリューを$\operatorname{Attention}$に送る前に、学習される線形変換行列によってそれぞれをあらかじめ複数の空間へ写す。
クエリ・キー・バリューベクトルの次元が$d_{\text{model}}$であるとしよう。言い換えれば$\mathbf{Q} \in M^{d_{\text{model}} \times n}$、$\mathbf{K} \in M^{d_{\text{model}} \times m}$、$\mathbf{V} \in M^{d_{\text{model}} \times m}$である。それぞれの$h = 1, \cdots, H$に対して、クエリ・キーを$d_{k}$次元へ、バリューを$d_{v}$次元へ圧縮する線形変換が次のように与えられたとしよう。
$$ \mathbf{W}_{h}^{Q} \in M^{d_{k} \times d_{\text{model}}}, \qquad \mathbf{W}_{h}^{K} \in M^{d_{k} \times d_{\text{model}}}, \qquad \mathbf{W}_{h}^{V} \in M^{d_{v} \times d_{\text{model}}} $$
すると$(\mathbf{Q}, \mathbf{K}, \mathbf{V})$はヘッドごとに互いに異なる空間へ写され、$H$種類の情報へ、すなわち$H$個のヘッドへ分解される。
$$ \operatorname{head}_{h} := \operatorname{Attention} (\mathbf{W}_{h}^{Q}\mathbf{Q}, \mathbf{W}_{h}^{K}\mathbf{K}, \mathbf{W}_{h}^{V}\mathbf{V}) \in M^{d_{v} \times n} \tag{2} $$
$(1)$では$\mathbf{Q} \in M^{d_{k} \times n}$であるのに対し、$(2)$では$\mathbf{Q} \in M^{d_{\text{model}} \times n}$であり$\mathbf{W}_{h}^{Q}\mathbf{Q} \in M^{d_{k} \times n}$であることに注意せよ。
$$ \begin{align*} \text{in } (1) &: \quad \mathbf{Q} \in M^{d_{k} \times n} \\ \text{in } (2) &: \quad \mathbf{Q} \in M^{d_{\text{model}} \times n}, \quad \mathbf{W}_{h}^{Q}\mathbf{Q} \in M^{d_{k} \times n} \end{align*} $$
これらをすべて集めると、以下のような形で表現できる。
$$ \begin{bmatrix} \operatorname{head}_{1} \\ \operatorname{head}_{2} \\ \vdots \\ \operatorname{head}_{H} \end{bmatrix} = \begin{bmatrix} \operatorname{Attention}(\mathbf{W}_{1}^{Q}\mathbf{Q}, \mathbf{W}_{1}^{K}\mathbf{K}, \mathbf{W}_{1}^{V}\mathbf{V}) \\ \operatorname{Attention}(\mathbf{W}_{2}^{Q}\mathbf{Q}, \mathbf{W}_{2}^{K}\mathbf{K}, \mathbf{W}_{2}^{V}\mathbf{V}) \\ \vdots \\ \operatorname{Attention}(\mathbf{W}_{H}^{Q}\mathbf{Q}, \mathbf{W}_{H}^{K}\mathbf{K}, \mathbf{W}_{H}^{V}\mathbf{V}) \end{bmatrix} \in M^{H d_{v} \times n} $$
そして次のようなブロック行列 $\mathbf{W}^{O} \in M^{d_{\text{model}} \times H d_{v}}$を掛ければ、ヘッドたちの線形結合を得ることができる。
$$ \mathbf{W}^{O} = \begin{bmatrix} \mathbf{W}_{1}^{O} & \mathbf{W}_{2}^{O} & \cdots & \mathbf{W}_{H}^{O} \end{bmatrix}, \quad \mathbf{W}_{h}^{O} \in M^{d_{\text{model}} \times d_{v}} $$
$$ \begin{bmatrix} \mathbf{W}_{1}^{O} & \mathbf{W}_{2}^{O} & \cdots & \mathbf{W}_{H}^{O} \end{bmatrix} \begin{bmatrix} \operatorname{head}_{1} \\ \operatorname{head}_{2} \\ \vdots \\ \operatorname{head}_{H} \end{bmatrix} = \sum_{h=1}^{H} \mathbf{W}_{h}^{O} \operatorname{head}_{h} \in M^{d_{\text{model}} \times n} $$
この行列がマルチヘッドアテンションの出力である。ここで関数$\operatorname{MultiHead}$を以下のように定義しよう。
定義
次元が$d_{\text{model}}$であるベクトルたちの行列$\mathbf{Q} \in M^{d_{\text{model}} \times n}$、$\mathbf{K} \in M^{d_{\text{model}} \times m}$、$\mathbf{V} \in M^{d_{\text{model}} \times m}$を、それぞれクエリ行列、キー行列、バリュー行列としよう。それぞれの$h = 1, \cdots, H$に対して、これらの行列の列ベクトルをそれぞれ$d_{k}$、$d_{k}$、$d_{v}$次元へ写す線形変換たちを$\mathbf{W}_{h}^{Q} \in M^{d_{k} \times d_{\text{model}}}$、$\mathbf{W}_{h}^{K} \in M^{d_{k} \times d_{\text{model}}}$、$\mathbf{W}_{h}^{V} \in M^{d_{v} \times d_{\text{model}}}$としよう。$\operatorname{head}_{h}$を以下のように定義しよう。
$$ \operatorname{head}_{h} := \operatorname{Attention} (\mathbf{W}_{h}^{Q}\mathbf{Q}, \mathbf{W}_{h}^{K}\mathbf{K}, \mathbf{W}_{h}^{V}\mathbf{V}) \in M^{d_{v} \times n} $$
関数マルチヘッドアテンションmulti-head attention $\operatorname{MultiHead} : M^{d_{\text{model}} \times n} \times M^{d_{\text{model}} \times m} \times M^{d_{\text{model}} \times m} \to M^{d_{\text{model}} \times n}$を以下のように定義する。
$$ \begin{align*} \operatorname{MultiHead} (\mathbf{Q}, \mathbf{K}, \mathbf{V}) &:= \mathbf{W}^{O} \begin{bmatrix} \operatorname{head}_{1} \\ \operatorname{head}_{2} \\ \vdots \\ \operatorname{head}_{H} \end{bmatrix} \\ &= \sum_{h=1}^{H} \mathbf{W}_{h}^{O} \operatorname{head}_{h} \\ &= \sum_{h=1}^{H} \mathbf{W}_{h}^{O} \operatorname{Attention} (\mathbf{W}_{h}^{Q}\mathbf{Q}, \mathbf{W}_{h}^{K}\mathbf{K}, \mathbf{W}_{h}^{V}\mathbf{V}) \\ &= \sum_{h=1}^{H} \mathbf{W}_{h}^{O} \left( \mathbf{W}_{h}^{V}\mathbf{V} \right) \operatorname{Softmax} \left( f(\mathbf{W}_{h}^{Q}\mathbf{Q}, \mathbf{W}_{h}^{K}\mathbf{K}) \right) \end{align*} $$
ここで$\mathbf{W}^{O} = \begin{bmatrix} \mathbf{W}_{1}^{O} & \cdots & \mathbf{W}_{H}^{O} \end{bmatrix} \in M^{d_{\text{model}} \times H d_{v}}$は学習される重み行列であり、各ブロックのサイズは$\mathbf{W}_{h}^{O} \in M^{d_{\text{model}} \times d_{v}}$である。
数学の論文において2 3
数学の論文で主に使われる定義は、上のものよりもはるかに単純である。セルフアテンションの場合を考えると、一つのデータ$\mathbf{X}$から次のようにクエリ、キー、バリューを得る。
$$ \mathbf{Q} = \mathbf{W}^{QX}\mathbf{X},\qquad \mathbf{K} = \mathbf{W}^{KX}\mathbf{X},\qquad \mathbf{V} = \mathbf{W}^{VX}\mathbf{X} $$
するとアテンションの入力が$\mathbf{W}^{X}\mathbf{W}^{QX}\mathbf{X}$のように煩雑な形で表現されるが、線形変換の合成もまた線形変換であるから、$\mathbf{W}^{X}\mathbf{X}$と表現するのと本質的に同じである。また、数学的にはアテンション関数はマルチヘッドアテンションにおいて$H=1$である特殊な場合に過ぎないので、残差層を含みスコア関数を内積とした次の関数をアテンションと定義する。
$$ \operatorname{Attention}(\mathbf{X}) = \mathbf{X} + \sum_{h=1}^{H} \left( \mathbf{W}_{h}^{V}\mathbf{X} \right) \operatorname{Softmax} \left[ (\mathbf{W}_{h}^{K}\mathbf{X})^{\mathsf{T}} (\mathbf{W}_{h}^{Q}\mathbf{X}) \right] $$
説明
トランスフォーマー論文では行ベクトルを基準に記述しているため行列たちの次元が逆であり、記法も以下のようになっている。
$$ \operatorname{head}_{h} = \operatorname{Attention} (\mathbf{Q}\mathbf{W}_{h}^{Q}, \mathbf{K}\mathbf{W}_{h}^{K}, \mathbf{V}\mathbf{W}_{h}^{V}) \in M^{n \times d_{v}} $$
$$ \operatorname{MultiHead} (\mathbf{Q}, \mathbf{K}, \mathbf{V}) = \begin{bmatrix} \operatorname{head}_{1} & \operatorname{head}_{2} & \cdots & \operatorname{head}_{H} \end{bmatrix} \begin{bmatrix} \mathbf{W}_{1}^{O} \\ \mathbf{W}_{2}^{O} \\ \vdots \\ \mathbf{W}_{H}^{O} \end{bmatrix} $$
Ashish Vaswani et al. Attention is all you need. Advances in neural information processing systems 30 (2017). ↩︎
Chulhee Yun et al. Are transformers universal approximators of sequence-to-sequence functions?. arXiv preprint arXiv:1912.10077 (2019). ↩︎
Silas Alberti et al. Sumformer: Universal approximation for efficient transformers. Topological, Algebraic and Geometric Learning Workshops 2023. PMLR, 2023. ↩︎
