logo

マスクドアテンション 📂機械学習

マスクドアテンション

定義

それぞれクエリ、キー、バリューと呼ばれる三つの行列$\mathbf{Q} \in M^{d_{k} \times n}$、$\mathbf{K} \in M^{d_{k} \times m}$、$\mathbf{V} \in M^{d_{v} \times m}$が与えられたとしよう。成分が$0$または$-\infty$である行列$\mathbf{M} \in M^{m \times n}$をマスクmaskといい、次のように定義される関数をマスクドアテンションmasked attentionという。

$$ \operatorname{MaskedAttention}(\mathbf{Q}, \mathbf{K}, \mathbf{V}) := \mathbf{V} \operatorname{Softmax} \left( \frac{\mathbf{K}^{\mathsf{T}} \mathbf{Q}}{\sqrt{d_{k}}} + \mathbf{M} \right) \tag{1} $$

ここで$\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} $$

説明

$[\mathbf{M}]_{ij} = 0$であるということは$\operatorname{score}(\mathbf{k}_{i}, \mathbf{q}_{j})$をそのままにしておくということであり、$[\mathbf{M}]_{ij} = -\infty$であるということは$\operatorname{score}(\mathbf{k}_{i}, \mathbf{q}_{j})$の値を補正して$-\infty$に強制するという意味である。ソフトマックスの定義により、$\operatorname{softmax}(-\infty) = 0$であるため、マスキングされた成分はソフトマックスを経てアテンション確率がちょうど$0$になる。すなわちマスキングとは、クエリが特定のキー/バリューを見られないように強制的に覆い隠す装置である。

トランスフォーマー論文では、デコーダの最初のサブ層でクエリが自分自身とそれ以前のキーだけを見られるようにする因果マスクcausal maskを使用する。すなわち$i$番目のクエリが$j$番目のキーを参照できるかどうかは、$i \gt j$であるかによって決まる。これを関数として定義すると次のとおりである。

$$ \begin{bmatrix} \operatorname{Mask}(\mathbf{K}^{\mathsf{T}} \mathbf{Q}/ \sqrt{d_{k}}) \end{bmatrix}_{ij} = \begin{cases} -\infty & i \gt j \\ \mathbf{k}_{i}^{\mathsf{T}} \mathbf{q}_{j} / \sqrt{d_{k}} & i \le j \end{cases} $$

$$ \operatorname{MaskedAttention}(\mathbf{Q}, \mathbf{K}, \mathbf{V}) := \mathbf{V} \operatorname{Softmax} \circ \operatorname{Mask} \left( \frac{\mathbf{K}^{\mathsf{T}} \mathbf{Q}}{\sqrt{d_{k}}} \right) $$

一方、実際の実装では$-\infty$の代わりに$-10^{9}$程度の非常に大きな負の数を加えるのが普通である。

関連リンク