LLMは文章を理解する際、どの単語が、どの単語と深く関係しているか、を計算する。この仕組みを Attention という。この Attention を行う回路(ヘッド)を複数(マルチ)並列で動かすのが マルチヘッド・アテンション(MHA)。
1つのアテンション回路(単一のドットプロダクト・アテンション)は、あるトークンから他のトークンへのアテンションの重み(分布)を1パターンしか出力できない。しかし、言語の文脈において、ある単語は他の複数の単語と異なる性質の関係(構文的関係、意味的関係、長距離の依存関係など)を同時に持っている。1つの重みベクトル(1パターンの分布)だけでは、これらの異なる関係性を同時に数値化(エンコード)することが不可能。
マルチヘッド・アテンション(MHA)では、元々の埋め込み次元($[d_{\text{model}}$)をヘッド数($h$)で割った、より小さな次元($d_k = d_{\text{model}} / h$)の空間を複数作る。 それぞれのヘッドが独立した異なる重み行列($W_Q^{(i)}, W_K^{(i)}, W_V^{(i)}$)を持つため、各ヘッドは全く異なるアテンション分布を独立して計算できる。
これらを並列で計算し、最後に各ヘッドの出力($V$ に重みを掛けたもの)を結合(Concatenate)して再度線形変換することにより、多様な特徴(関係性)を1つのベクトルに同時に内包させて次の層へ送ることができる。
つまり、表現力(表現容量)を制限させずに、異なる幾何学的空間における相関関係を同時に捉えるために、複数のヘッドが必要となる。
ここからは、MHAが実際にどう計算されるかを、PyTorchのサンプルコードと対応させて順に見る。入力は $X \in \mathbb{R}^{n \times d_{\text{model}}}$ である。コードでは記号が異なる。
| 本文の記号 | コード上の名前 | 意味 |
|---|---|---|
| $n$ | C(context_len) | 系列長 |
| $d_{\text{model}}$ | E(embed_dim) | 埋め込み次元 |
| $h$ | H(n_head) | ヘッド数 |
| $d_k$ | D(head_dim) | ヘッド次元 |
| (なし) | B | バッチサイズ |
コードの head_dim は独立した引数である。$d_k = d_{\text{model}}/h$ を強制しない。$h \cdot d_k \ne d_{\text{model}}$ でも動く。その場合、後述の $W_O$ の形は $\mathbb{R}^{hd_k \times d_{\text{model}}}$ になる。
ヘッド $i$ ごとに、入力 $X$ から3種類のベクトルを作る。\[Q_i = X W_Q^{(i)},\quad K_i = X W_K^{(i)},\quad V_i = X W_V^{(i)}\]
各重み行列の形は $W^{(i)} \in \mathbb{R}^{d_{\text{model}} \times d_k}$ である。役割は次のとおり。
実装では、$h$ 個の重み行列を横に連結した1つの大きな行列にまとめる。\[W_Q = \left[ W_Q^{(1)} \mid \cdots \mid W_Q^{(h)} \right] \in \mathbb{R}^{d_{\text{model}} \times hd_k}\]
これが nn.Linear(E, H * D, bias=False) である。1回の行列積で全ヘッド分を計算できる。GPUでの並列効率が高い。$W_K$、$W_V$ も同様である。
射影の出力は (B, C, H*D) である。これを次の2段階で (B, H, C, D) に変える。
view(B, C, H, D): 最後の軸を、ヘッド軸とヘッド内の次元軸に分けるtranspose(1, 2): ヘッド軸を前に出すヘッド軸をバッチ軸の隣に置くと、以降の matmul が B と H を独立に扱う。全ヘッドが並列に計算される。
クエリとキーの内積で、トークン間の関連度を測る。\[S_i = \frac{Q_i K_i^\top}{\sqrt{d_k}} \in \mathbb{R}^{n \times n}\]
$S_i$ の $(p, q)$ 成分は、「トークン $p$ がトークン $q$ をどれだけ見るか」を表す。
コードでは、K.transpose(-2, -1) で (B, H, D, C) にして matmul する。結果は (B, H, C, C) になる。
$\sqrt{d_k}$ で割る理由は次のとおり。成分が平均0・分散1の独立な変数なら、内積の分散は $d_k$ になる。割らないと、$d_k$ が大きいほどスコアが極端な値になる。softmax が飽和し、勾配がほぼ0になる。$\sqrt{d_k}$ で割ると、分散が1に戻る。
文章を先頭から生成するLLMでは、トークンは未来のトークンを見てはいけない。そこで、未来の位置に $-\infty$ を加える。\[M_{pq} =\begin{cases}0 & (q \le p) \\-\infty & (q > p)\end{cases}\qquad\tilde{S}_i = S_i + M\]
コードでは、下三角行列 torch.tril でマスクを作る。0の位置を masked_fill で -inf にする。$n = 4$ の場合、マスクは次の形である。1が見てよい位置、0が隠す位置。\[\begin{bmatrix}1 & 0 & 0 & 0 \\1 & 1 & 0 & 0 \\1 & 1 & 1 & 0 \\1 & 1 & 1 & 1\end{bmatrix}\]
対角成分は必ず1である。どの行も全部が $-\infty$ にはならない。softmax が NaN にならない。
マスクは register_buffer で登録している。
.to(device) でモデルと一緒にGPUへ移動するpersistent=False により、state_dict に保存されないforward のたびにマスクを作り直す必要がない。使うときは self.mask[:C, :C] で必要な大きさだけ切り出す。
行ごとに softmax を取り、和が1の分布にする。\[A_i = \mathrm{softmax}(\tilde{S}_i)\]
$e^{-\infty} = 0$ なので、マスクされた位置の重みは0になる。この $A_i$ が、冒頭で述べた「アテンションの重み(分布)」である。ヘッドごとに別の $A_i$ ができる。
コードでは、この $A_i$ に Dropout(attention_dropout)をかける。特定のトークンへの過度な依存を防ぐ正則化である。
重みを使って $V$ を混ぜ合わせ、各ヘッドの出力を得る。\[\mathrm{head}_i = A_i V_i \in \mathbb{R}^{n \times d_k}\]
関係の深いトークンの $V$ ほど、大きく取り込まれる。コード上の形は、(B, H, C, C) と (B, H, C, D) の積で (B, H, C, D) になる。
全ヘッドの出力を横に並べ、線形変換で元の次元に戻す。\[\mathrm{MHA}(X) = \mathrm{Concat}(\mathrm{head}_1, \dots, \mathrm{head}_h)\, W_O,\qquad W_O \in \mathbb{R}^{d_{\text{model}} \times d_{\text{model}}}\]
連結後の次元は $h \cdot d_k = d_{\text{model}}$ である。入力と出力の形が同じなので、残差接続でそのまま足し合わせられる。$W_O$ には、独立に計算された各ヘッドの情報を混ぜ合わせる役割がある。
コードでは、連結を次の順で行う。
transpose(1, 2) で (B, C, H, D) に戻すcontiguous() でメモリ配置を整えるview(B, C, H * D) でヘッド軸を連結するW_o で (B, C, E) に射影するtranspose 後のテンソルは、メモリ上で連続していない。連続していないテンソルには view が使えない。そのため contiguous() が必要になる。
最終出力にも Dropout(output_dropout)をかける。Dropout は学習時だけ有効で、eval() で無効になる。
ヘッドを増やしても、計算量はほとんど増えない。$h$ 個のヘッドのスコア計算の量は、次のとおりである。\[h \cdot n^2 \cdot d_k = n^2 \cdot d_{\text{model}}\]
これは、次元 $d_{\text{model}}$ の単一ヘッドと同じ量である。つまり、同じ計算コストで、1パターンではなく $h$ パターンの分布を得られる。
パラメータ数も同様に、ヘッド数によらない。\[\underbrace{3 \cdot d_{\text{model}}^2}_{W_Q,\,W_K,\,W_V} + \underbrace{d_{\text{model}}^2}_{W_O} = 4\,d_{\text{model}}^2\]
上で挙げた「構文」「意味」「長距離依存」は、分かりやすくした例である。実際には、役割を人間が割り当てるわけではない。学習の過程で、各ヘッドが異なる関係を担うように分化していく。役割がはっきり読み取れるヘッドもあれば、そうでないヘッドもある。
スコア行列 $S_i$ の大きさは $n \times n$ である。系列長 $n$ の2乗で、計算量もメモリも増える。長い文章でAttentionが重くなる主因である。
サンプルコードは、実装の結果を F.scaled_dot_product_attention と比較している。この関数は、スケーリング・因果マスク・softmax・重み付き和を1つにまとめたものである。
環境によっては FlashAttention などの最適化カーネルが使われる。スコア行列 $n \times n$ をメモリに実体化せずに済む。手書き実装は仕組みの理解に向いている。実運用では標準関数のほうが速く、省メモリである。
比較の前に eval() を呼ぶのは、Dropout を無効にするためである。有効なままだと、出力が確率的に変わり一致しない。
| 段階 | 形状 | 数式 |
|---|---|---|
| 入力 | (B, C, E) | $X$ |
| Q, K, V 射影 | (B, C, H*D) | $XW_Q,\ XW_K,\ XW_V$ |
| ヘッド分割 | (B, H, C, D) | $Q_i,\ K_i,\ V_i$ |
| スコア | (B, H, C, C) | $Q_iK_i^\top / \sqrt{d_k} + M$ |
| Attention重み | (B, H, C, C) | $\mathrm{softmax}(\cdot)$ |
| ヘッド出力 | (B, H, C, D) | $A_iV_i$ |
| 連結 | (B, C, H*D) | $\mathrm{Concat}$ |
| 出力射影 | (B, C, E) | $\cdot\, W_O$ |
import torch
import torch.nn as nn
import torch.nn.functional as F
class MultiHeadAttention(nn.Module):
def __init__(self, embed_dim, n_head, head_dim, max_len=1024, dropout_rate=0.1):
super().__init__()
self.n_head = n_head
self.head_dim = head_dim
E, H, D = embed_dim, n_head, head_dim
self.W_q = nn.Linear(E, H * D, bias=False)
self.W_k = nn.Linear(E, H * D, bias=False)
self.W_v = nn.Linear(E, H * D, bias=False)
self.W_o = nn.Linear(H * D, E, bias=False)
self.attention_dropout = nn.Dropout(dropout_rate)
self.output_dropout = nn.Dropout(dropout_rate)
mask = torch.tril(torch.ones(max_len, max_len))
self.register_buffer("mask", mask, persistent=False)
def forward(self, x):
B, C, E = x.shape
H, D = self.n_head, self.head_dim
# Q, K, V を計算し、ヘッドごとに分割: (B, H, C, D)
Q = self.W_q(x).view(B, C, H, D).transpose(1, 2)
K = self.W_k(x).view(B, C, H, D).transpose(1, 2)
V = self.W_v(x).view(B, C, H, D).transpose(1, 2)
# スコア: (B, H, C, C)
scores = torch.matmul(Q, K.transpose(-2, -1)) / (D ** 0.5)
scores = scores.masked_fill(self.mask[:C, :C] == 0, float("-inf"))
# Attention重み → 重み付き和: (B, H, C, D)
weights = F.softmax(scores, dim=-1)
weights = self.attention_dropout(weights)
hidden = torch.matmul(weights, V)
# ヘッド結合: (B, C, H*D) → 出力射影: (B, C, E)
hidden = hidden.transpose(1, 2).contiguous().view(B, C, H * D)
return self.output_dropout(self.W_o(hidden))
# 動作確認
device = "cuda" if torch.cuda.is_available() else "cpu"
mha = MultiHeadAttention(embed_dim=512, n_head=8, head_dim=64).to(device)
x = torch.randn(2, 10, 512, device=device)
y = mha(x)
print(f"入力形状: {x.shape}") # (2, 10, 512)
print(f"出力形状: {y.shape}") # (2, 10, 512)
# PyTorch標準実装との一致確認
mha.eval()
with torch.no_grad():
B, C, _ = x.shape
H, D = mha.n_head, mha.head_dim
q = mha.W_q(x).view(B, C, H, D).transpose(1, 2)
k = mha.W_k(x).view(B, C, H, D).transpose(1, 2)
v = mha.W_v(x).view(B, C, H, D).transpose(1, 2)
ref = F.scaled_dot_product_attention(q, k, v, is_causal=True)
ref = mha.W_o(ref.transpose(1, 2).contiguous().view(B, C, H * D))
print("標準実装と一致:", torch.allclose(mha(x), ref, atol=1e-5))
Dropout は Attention重み $A_i$ の要素を、学習時にランダムに0にする層。nn.Dropout(p) が、確率 $p$ で各要素を0にする。
各要素に独立なマスク変数を $m_{pq}$ とすると、\[m_{pq} \sim \mathrm{Bernoulli}(1-p)\]
Dropout 後の重みは次のとおり。\[\tilde{A}_{pq} = \frac{m_{pq}}{1-p}\, A_{pq}\]
場合分けで書くと、次のようになる。\[\tilde{A}_{pq} =\begin{cases}0 & \text{確率 } p \\[4pt]\dfrac{A_{pq}}{1-p} & \text{確率 } 1-p\end{cases}\]
$1/(1-p)$ 倍して、残した要素を拡大するのは、出力の期待値を Dropout なしの場合と一致させるため。\[\mathbb{E}\left[\tilde{A}_{pq}\right]= \frac{\mathbb{E}[m_{pq}]}{1-p}\, A_{pq}= \frac{1-p}{1-p}\, A_{pq}= A_{pq}\]
この方式は inverted dropout と呼ばれる。推論時に特別な補正が要らなくなる。
softmax 直後は、各行の和が1となる。\[\sum_q A_{pq} = 1\]
Dropout 後は、和が1とは限らない。\[\sum_q \tilde{A}_{pq} \neq 1 \quad (\text{一般には})\]ただし期待値では1になる。$$\mathbb{E}\left[\sum_q \tilde{A}_{pq}\right] = \sum_q A_{pq} = 1$$
Attention の出力は、次の順で計算する。\[\mathrm{head}_i = \mathrm{Dropout}\!\left(\mathrm{softmax}(\tilde{S}_i)\right) V_i\]
softmax の後に適用する理由は、softmax の前に0を入れると確率として意味を持たなくなるため。例えば、スコアを0にすると、$e^0 = 1$ という有限の重みが残る。これは「そのトークンを見ない」ことにならない。softmax 後に0にすれば、そのトークンは確実に無視される。
因果マスクとの違いにも注意が必要。マスクは softmax の前に $-\infty$ を加える。一方、Dropout は softmax の後に0を掛ける。
Dropoutの効果は、
Mathematics is the language with which God has written the universe.