Self-Attention機構(Q, K, Vの基本)

トランスフォーマー(Transformer)の核心をなす Self-Attention(自己注意)機構は、入力シーケンス内の各要素が他のすべての要素との相互関係を直接評価し、動的に情報の重み付けを行う機構である。従来の RNN や CNN が持つ逐次性や局所性の制約を排除し、全結合的な依存関係を一度に処理できる点に本質的な特徴がある。この構造により、長距離依存関係のモデリングと高い並列計算効率が同時に達成される。

ただし重要な注意点として、Self-Attention は順序不変(permutation-equivariant)である。すなわち、入力トークンの順序を入れ替えても同じ注意重みが計算される。このため、実際の Transformer では Positional Encoding(位置符号化)を入力埋め込みに加算し、系列の位置情報を明示的に注入する必要がある。

基本設定

入力シーケンスを $n$ 個のトークンからなる行列 $X \in \mathbb{R}^{n \times d}$ とする。ここで $n$ は系列長、$d$ は埋め込み次元である。Self-Attention では、この入力に対して3つの線形変換を施し、Query(クエリ)、Key(キー)、Value(バリュー)を生成する。

\[Q = XW^Q, \quad K = XW^K, \quad V = XW^V\]

ここで $W^Q, W^K \in \mathbb{R}^{d \times d_k}$、$W^V \in \mathbb{R}^{d \times d_v}$ は学習可能なパラメータ行列である。シングルヘッドの場合は $d_k = d_v = d$ とすることも可能だが、Multi-Head Attention では通常 $d_k = d_v = d/h$($h$ はヘッド数)と設定される。

  • Query($Q$): 各トークンが「何を求めているか」を表す特徴ベクトル。
  • Key($K$): 各トークンが持つ情報の「検索キー」であり、Query との適合度を測る。
  • Value($V$): 実際に集約される情報本体であり、Attention によって重み付けされる。

Scaled Dot-Product Attention の数理

Self-Attention の計算は、Query と Key の内積に基づく類似度計算から始まる。

\[\text{Attention}(Q, K, V) = \text{softmax}\!\left(\frac{QK^T}{\sqrt{d_k}}\right)V\]

ここで $QK^T \in \mathbb{R}^{n \times n}$ はトークン間の相関スコア行列である。

スケーリング因子の厳密な意味

$q_l, k_l$ を平均 0、分散 1 の独立確率変数とすると、

\[\mathrm{Var}(q_l k_l) = 1\]\[\mathrm{Var}(q \cdot k) = \sum_{l=1}^{d_k} \mathrm{Var}(q_l k_l) = d_k\]

したがってスケーリングなしでは softmax が飽和するため、

\[\mathrm{Var}\!\left(\frac{q \cdot k}{\sqrt{d_k}}\right) = 1\]

となるように正規化する。

softmax の役割

\[\alpha_{ij} =\frac{\exp\!\left(\dfrac{q_i \cdot k_j}{\sqrt{d_k}}\right)}{\sum_{j'=1}^{n} \exp\!\left(\dfrac{q_i \cdot k_{j'}}{\sqrt{d_k}}\right)}\]\[\mathrm{output}_i = \sum_{j=1}^{n} \alpha_{ij} v_j\]

これにより、各クエリごとに確率分布が形成され、Value の加重平均が計算される。

Causal Mask(因果マスク)

\[\tilde{S}_{ij} =\begin{cases}\dfrac{q_i \cdot k_j}{\sqrt{d_k}} & (j \leq i) \\-\infty & (j > i)\end{cases}\]\[A = \text{softmax}(\tilde{S})\]

これにより未来トークンへの依存が遮断される。

Multi-Head Attention への拡張

\[\text{head}_i = \text{Attention}(QW_i^Q,\; KW_i^K,\; VW_i^V)\]\[\text{MultiHead}(Q, K, V)= \text{Concat}(\text{head}_1, \ldots, \text{head}_h)\, W^O\]

各ヘッドは異なる部分空間で相関を学習し、多様な依存関係を同時に表現する。

計算量と構造的性質

Self-Attention の計算量は、主としてスコア行列 $QK^T \in \mathbb{R}^{n \times n}$ の計算に支配される。この行列は、すべてのトークン対 $(i, j)$ に対する内積 $q_i \cdot k_j$ を計算することで得られるため、計算量は以下のようになる。

\[\text{時間計算量:}\mathcal{O}(n^2 d)\]

ここで $n^2$ はトークン間の全結合的な相互作用の数を表し、各内積計算に $d$ 次元の演算が必要となるため、全体として $\mathcal{O}(n^2 d)$ となる。この二乗スケーリングは、長い系列に対して計算コストが急激に増大することを意味し、Self-Attention の最大のボトルネックである。

また、softmax によって得られる注意重み行列 $A \in \mathbb{R}^{n \times n}$ を保持する必要があるため、空間計算量(メモリ使用量)も同様に次式で与えられる。

\[\text{空間計算量:}\mathcal{O}(n^2)\]

このメモリ使用量は特に GPU 上での実装において重要な制約となり、大規模モデルや長文処理においてはバッチサイズの制限や分割計算を余儀なくされる要因となる。

一方で、Self-Attention は極めて重要な構造的利点を持つ。それが「最大経路長(maximum path length)」の短さである。任意の2つのトークン間の情報伝播に必要な計算ステップ数は次のように評価される。

\[\text{最大経路長:}\mathcal{O}(1)\]

これは、Self-Attention ではすべてのトークンが1回の演算で直接相互作用できるためである。すなわち、トークン $i$ の情報は、単一の Attention 層を通じて即座にトークン $j$ に伝播する。この性質は、長距離依存関係の学習において決定的に重要である。

これに対して、他の代表的な系列モデルと比較すると次のような違いがある。

  • RNN: 情報は逐次的に伝播するため、トークン間距離に比例して経路長が増加する。\[\text{最大経路長(RNN):}\mathcal{O}(n)\]長距離依存関係では勾配消失・爆発が起きやすい。
  • CNN: カーネルサイズ $k$ に依存した局所的な受容野を持つ。\[\text{最大経路長(CNN):}\mathcal{O}(\log_k n)\]多層化により受容野は拡大するが、完全な全結合関係にはならない。

さらに並列計算性の観点から見ると、Self-Attention は全トークンを同時に処理できるため、高い並列化効率を持つ。一方で RNN は逐次依存性のため並列化が困難であり、CNN は局所的には並列化可能だが受容野の制約を受ける。

以上を整理すると、Self-Attention は

  • 全結合的相互作用による高い表現力
  • 最大経路長 $\mathcal{O}(1)$ による優れた長距離依存性の捕捉能力
  • 高い並列計算効率
を持つ一方で、
  • 時間計算量 $\mathcal{O}(n^2 d)$
  • 空間計算量 $\mathcal{O}(n^2)$
というスケーリング上の課題を抱える。このトレードオフが、Sparse Attention、Linear Attention、FlashAttention などの効率化手法が研究される根本的な動機となっている。

モデル時間計算量最大経路長並列計算
Self-Attention$\mathcal{O}(n^2 d)$$\mathcal{O}(1)$○
RNN$\mathcal{O}(nd^2)$$\mathcal{O}(n)$✗
CNN$\mathcal{O}(knd^2)$$\mathcal{O}(\log_k n)$△

まとめ

Self-Attention は Q・K・V の三分解に基づき、系列全体の相互作用を一度に計算する機構である。スケーリング付き内積と softmax による確率的重み付け、Multi-Head による多視点表現により極めて高い表現能力を実現する。一方で計算量が $\mathcal{O}(n^2)$ にスケールするため、これが各種効率化手法の出発点となっている。

Mathematics is the language with which God has written the universe.





















数理統計学 機械学習