アテンションは加重平均だ
一言でいうと
アテンションは、「いまこの位置で、どの位置をどれだけ見るか」を決めて、値の加重平均をとる演算です。それがすべてです。
3つの名前
同じ入力から、3つのベクトルを作ります。
- Q(query): 探しているもの
- K(key): 各位置が掲げる目印
- V(value): その位置が実際に差し出す内容
図書館のたとえがぴったりです。質問(Q)を持って行き、背表紙のタイトル(K)と照らし合わせ、よく合う本の内容(V)を、その度合いに応じて持ち帰ります。
score = Q · Kᵀ / √d ← 얼마나 맞는가
weight = softmax(score) ← 합이 1이 되게
out = weight · V ← 가중 평균
たった3行です。残りはすべて、この3行を何度も、いろいろな方向に行うことです。
なぜ√dで割るのか、数字で確かめる
「そういうものだ」で済まされることが最も多い部分ですが、理由ははっきりしています。
成分が独立で分散が1である2つのd次元ベクトルの内積は、分散がdになります。d=64なら、スコアはおよそ±8の範囲に散らばります。
ソフトマックスに±8のスコアが入ると、どうなるでしょうか。比はe⁸ / e⁻⁸ ≈ 9백만(韓国語で「900万」を意味する語です)になります。最大の1つがほぼ1をすべて持っていき、残りは0になります。加重平均ではなく、単なる「1つ選び」になってしまいます。
そうなると、学習ができません。ソフトマックスが飽和すると、勾配が0に近づくからです。
√dで割ると、分散が再び1に戻ります。このラボのステップ2で、これを実際に測ります。割らなかったときに、最大の確率がどこまで上がるかを見ます。
最大値を引いて計算するソフトマックス
def softmax(xs):
m = max(xs) # 이 한 줄이 없으면
e = [exp(x - m) for x in xs] # exp(1000) 에서 터진다
s = sum(e)
return [v / s for v in e]
exp(x - max)は、数学的には同じ値ですが、オーバーフローが起きません。実務のコードでこの1行を抜かすと、大きな値が入った瞬間にinfとnanが出ます。
マスク、未来を見せないために
言語モデルは、次のトークンを当てるように学習します。ところがアテンションは、デフォルトではすべての位置を見ます。答えを見て答えを当てるようなものです。
そこで、スコア行列の上三角を-infにします。
k0 k1 k2
q0 0.3 -inf -inf
q1 0.1 0.5 -inf
q2 0.2 0.1 0.4
-infはexpを通ると、ちょうど0になります。0を掛けるのではなく、ソフトマックスの前に-infを足すことが核心です。ソフトマックスのあとに0を掛けると、残りの重みの合計が1にならなくなります。
これが「causal」または「decoder」マスクです。BERTのようなエンコーダーモデルにはありません。そのためBERTは文全体を見て理解することに強く、GPTは続きを書くことに強いのです。
マルチヘッド、なぜ分けるのか
d=64を一度に使う代わりに、8個に分けてd=8のアテンションを8回行い、そのあと連結します。計算量はほぼ同じです。
なぜでしょうか。1つのソフトマックスは、1種類の関係しか表現できません。重みの合計が1なので、複数の場所を同時に強く見ることはできません。ヘッドを分けると、あるヘッドは直前の単語を、あるヘッドは文頭の主語を、あるヘッドは引用符の対応を見ます。
学習済みのモデルのヘッドを開いてみると、実際にそのように分かれています。
位置エンコーディング、順序を知らないアテンション
これは、初めて学ぶときに最も驚く部分です。
アテンションには、順序という概念がまったくありません。入力トークンを並べ替えると、出力も同じように並べ替えられて出てくるだけで、値そのものは変わりません(permutation equivariant)。「私はあなたが好きです」と「あなたを私は好きです」が区別されません。
そこで、位置の情報を入力に足してあげます。
- サイン/コサイン(原論文): 学習パラメーターがなく、長い長さへの外挿ができます
- 学習型の埋め込み(BERT、ViT): 単純ですが、学習時に見た長さを超えられません
- RoPE(LLaMA、Qwenなど、最近のほとんど): Q・Kを回転させて、相対位置を自然に含めます。長さの拡張が容易なので、コンテキストを延ばしている最近のモデルは、すべてこれを使っています
- ALiBi: 距離に比例したペナルティを、スコアに足します
ステップ5で、これを実際に確認します。位置エンコーディングなしで入力を並べ替えると、出力がそのまま並べ替えられ、足すとそうなりません。
LayerNormと残差接続
x = x + Attention(LayerNorm(x))
x = x + FFN(LayerNorm(x))
- 残差(residual): 深く積み重ねても、勾配が入力まで届きます。これがなければ、6層を超えただけで学習できなくなります
- LayerNorm: 各トークンのベクトルを、平均0、分散1に揃えます。バッチではなく特徴軸で正規化する点がBatchNormとの違いで、そのため、バッチサイズやシーケンス長の影響を受けません
LayerNormをアテンションの前に置くか後ろに置くか(pre-LN vs post-LN)は、実際に学習の安定性を左右します。原論文はpost-LNでしたが、最近はほとんどがpre-LNです。ウォームアップなしでも学習できるからです。
コストはどこから生まれるのか
スコア行列はn × nです。シーケンス長の2乗でメモリと計算が増えます。コンテキストを4kから8kに延ばすと、4倍です。
これを減らそうとする試みが、最近の研究の大きな流れです。
- FlashAttention: 数学はそのままに、GPUのメモリアクセスの順序を変えて、実際の速度を数倍に引き上げます。近似ではありません
- GQA / MQA: K・Vヘッドを複数のQヘッドで共有して、推論時のKVキャッシュを減らします。最近のオープンモデルの大半がGQAです
- MoE: 層ごとに複数いる専門家のうち一部だけを有効にして、パラメーターは大きく、計算は小さく保ちます
- スライディングウィンドウ / スパースアテンション: 遠い位置は、そもそも見ません
まとめ
アテンションそのものは3行です。残りは、その3行を、安定して、安く、順序を知りながら行う方法についてのことです。次のラボで、その3行を自分で書き、スケーリングと位置エンコーディングがなければ何が崩れるのかを、数字で見ます。
現場では
この計算を自分で組む機会は、ほとんどありません。フレームワークがすべてやってくれるからです。それでも知っておくべき理由は、問題が起きたときにどこを見るかが、ここで分かれるからです。
学習の損失が突然nanになったら、たいていはソフトマックスの前の値があふれており、マスクを間違ってかけると、エラーなしに、静かに未来のトークンを見ながら学習して、評価スコアだけが妙に良く出ます。バッチサイズを変えたのに結果が変わるなら、正規化の軸を疑う必要があり、コンテキストを2倍にしたらメモリが4倍に跳ね上がるのは、バグではなく、構造上当然のことです。
推論サービスを運用するなら、KVキャッシュがそのままメモリの予算です。同時リクエスト数と最大コンテキスト長を掛けた値がGPUメモリに収まるかをまず計算し、収まらなければ、GQAを使うモデルを選ぶか、コンテキストの上限を下げるのが、実際の選択肢です。