マルチヘッド・アテンション

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}}}$ になる。

Q・K・V への射影

ヘッド $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) に変える。

  1. view(B, C, H, D): 最後の軸を、ヘッド軸とヘッド内の次元軸に分ける
  2. 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 で登録している。

forward のたびにマスクを作り直す必要がない。使うときは self.mask[:C, :C] で必要な大きさだけ切り出す。

softmax による重みの正規化

行ごとに 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$ には、独立に計算された各ヘッドの情報を混ぜ合わせる役割がある。

コードでは、連結を次の順で行う。

  1. transpose(1, 2) で (B, C, H, D) に戻す
  2. contiguous() でメモリ配置を整える
  3. view(B, C, H * D) でヘッド軸を連結する
  4. W_o で (B, C, E) に射影する

transpose 後のテンソルは、メモリ上で連続していない。連続していないテンソルには view が使えない。そのため contiguous() が必要になる。

最終出力にも Dropout(output_dropout)をかける。Dropout は学習時だけ有効で、eval() で無効になる。

なぜ $d_k = d_{\text{model}}/h$ に分けるのか

ヘッドを増やしても、計算量はほとんど増えない。$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

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の効果は、


@2026-09-29

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





















ELFプログラムヘッダーのデータ構造と型 学習式位置エンコーディング Dockerのネットワーク構成 kindのインストール Network Namespaceによる最小ネットワーク構成 固定式位置エンコーディング