一文でまとめると:MHA、MQA、GQA は一つのパラメータ族です。KV ヘッド数を変えると、品質・メモリ・速度のトレードオフ上を移動します。最適点は実際のモデルとハードウェアで測らなければ分かりません。
23.1 問題は「テンソルが何個あるか」ではない
第22章で見たように、新しい token を生成するとき、各レイヤーは新しい Q、K、V を計算します。Q はその場で使い終わりますが、新しい K と V はそのレイヤーの KV キャッシュに追加されます。
ここにはよくある数え間違いがあります。32 ヘッドだからといって、レイヤーごとに 64 個の独立したキャッシュテンソルを持つわけではありません。一般的な実装では、論理的には次の二つです。
K cache: [B, H_KV, S, d_h]
V cache: [B, H_KV, S, d_h]
ヘッド数 はテンソルの一つの軸です。保存量を決めるのは、この軸の長さです。
ここで は batch size、 はレイヤー数、 はキャッシュ済み token 数、 は head dimension、 は要素あたりの byte 数です。2 は K と V を表します。
この三つは、二つの数字だけでまとめられます。
- :Query ヘッド数
- :Key ヘッド数および Value ヘッド数
は で割り切れる必要があります。各 K/V ペアを共有する Query ヘッド数は次の通りです。
したがって:
| 方式 | 条件 | Query 8 ヘッドの例 |
|---|---|---|
| MHA | KV 8 ヘッド、1 group に Q 1 ヘッド | |
| GQA | KV 2 ヘッド、1 group に Q 4 ヘッド | |
| MQA | KV 1 ヘッドを Q 8 ヘッドで共有 |
MHA と MQA は別々の発明ではなく、GQA もその間に無理やり挟んだ第四の仕組みではありません。一つのパラメータ族の両端と、その間の領域です。
23.2 直感:八人の蒔絵師と図案帳
八人の蒔絵師が、修復途中の硯箱を調べているとします。ある人は螺鈿の欠けを見つけたい。別の人は梨地粉の粒度を確かめたい。さらに別の人は蓋の反りを調べたい。
- MHA:八人がそれぞれ自分の図案帳を持ち、異なる問いを立て、異なる方法で硯箱の特徴を記録します。
- MQA:問いは八通りのままですが、全員が一冊の共通図案帳を参照します。
- GQA:職人をいくつかの組に分け、組ごとに一冊の図案帳を共有します。別の組には別の帳面があります。
問い方が Q projection、過去を索引して持ち運ぶ方法が K と V に対応します。K/V を共有すると保存する履歴は減りますが、Query ヘッドがすべて同じになるわけではありません。Q の行と attention 出力は別々のままです。
三方式の形状は同じ記法で表せます。
q: [B, H_Q, T_q, d_h]
k: [B, H_KV, T_k, d_h]
v: [B, H_KV, T_k, d_h]
、 なら、対応は次の通りです。
Q0 Q1 Q2 Q3 -> KV0
Q4 Q5 Q6 Q7 -> KV1
概念上は、各 KV ヘッドを 回繰り返してから通常の MHA を計算できます。これは説明や参照実装には便利です。ただし効率のよい kernel は、キャッシュを実際に膨らませず、Query ヘッドを対応する KV ヘッドへ直接割り当てます。
23.3 何がどれだけ小さくなるのか
23.3.1 KV キャッシュ
第22章と同じ条件を使います。32 レイヤー、、1,024 tokens、FP16 または BF16(1 要素 2 bytes)、batch size 1 です。
| 方式 | 正確な byte 数 | 二進単位 | MHA 比 | |
|---|---|---|---|---|
| MHA | 32 | 536,870,912 | 512 MiB | 100% |
| GQA | 8 | 134,217,728 | 128 MiB | 25% |
| MQA | 1 | 16,777,216 | 16 MiB | 3.125% |
4,096 tokens では、それぞれ 2 GiB、512 MiB、64 MiB です。この比率自体は正確です。しかし「同時接続数が必ず 4 倍」「context が必ず 32 倍」とは言えません。server には weights、workspace、activation、allocator の余裕、断片化した block もあります。利用可能な context 長は positional encoding と学習時の長さにも制約されます。
23.3.2 Projection パラメータ
Bias を無視し、model width を とします。Q と output projection はそれぞれ 、K と V はそれぞれ 個の weight を持ちます。
の場合:
| 方式 | Attention projection weights |
|---|---|
| MHA、 | |
| GQA、 | |
| MQA、 |
32/32 の MHA から 32/8 の GQA に変えると、attention projection weights は 37.5% 減ります。7B モデル全体が 37.5% 減るわけでも、常に全体の 6% が減るわけでもありません。全体比率は FFN、embedding、depth、weight sharing によって変わります。
23.4 一つの実装で MHA、GQA、MQA を扱う
次のコードは一貫して [B, H, T, d_h] を使います。grouped_attention_reference は K と V を明示的に繰り返すため、動作を追いやすい参照実装です。grouped_attention は PyTorch の native enable_gqa 経路を使います。
import torch
import torch.nn.functional as F
def repeat_kv(x, num_query_heads):
"""Reference expansion: [B, H_kv, T, D_h] -> [B, H_q, T, D_h]."""
num_kv_heads = x.size(1)
if num_query_heads <= 0 or num_kv_heads <= 0:
raise ValueError("head counts must be positive")
if num_query_heads % num_kv_heads != 0:
raise ValueError("num_query_heads must be divisible by num_kv_heads")
return x.repeat_interleave(num_query_heads // num_kv_heads, dim=1)
def validate_gqa_shapes(q, k, v):
if q.ndim != 4 or k.ndim != 4 or v.ndim != 4:
raise ValueError("q, k, and v must be four-dimensional")
if q.size(0) != k.size(0) or q.size(0) != v.size(0):
raise ValueError("q, k, and v must have the same batch size")
if q.size(1) <= 0 or k.size(1) <= 0:
raise ValueError("head counts must be positive")
if q.size(1) % k.size(1) != 0 or k.size(1) != v.size(1):
raise ValueError("invalid GQA head counts")
if q.size(-1) != k.size(-1):
raise ValueError("q and k must have the same head dimension")
if k.size(-2) != v.size(-2):
raise ValueError("k and v must have the same sequence length")
def grouped_attention_reference(q, k, v, *, is_causal):
validate_gqa_shapes(q, k, v)
k_full = repeat_kv(k, q.size(1))
v_full = repeat_kv(v, q.size(1))
return F.scaled_dot_product_attention(
q, k_full, v_full, is_causal=is_causal, dropout_p=0.0
)
def grouped_attention(q, k, v, *, is_causal):
validate_gqa_shapes(q, k, v)
return F.scaled_dot_product_attention(
q,
k,
v,
is_causal=is_causal,
dropout_p=0.0,
enable_gqa=q.size(1) != k.size(1),
)
同じ関数の中で:
qが 8 ヘッド、k/vも 8 ヘッドなら MHA。qが 8 ヘッド、k/vが 2 ヘッドなら GQA。qが 8 ヘッド、k/vが 1 ヘッドなら MQA。
is_causal に既定値を置かず、必ず指定させているのは、Prefill とキャッシュ付き Decode で同じ値を使えるとは限らないからです。
- 学習時や、Q と K/V の長さが等しい Prefill では
is_causal=Trueを渡す。 - 1 token ずつの Decode で、K/V が過去と現在の token だけを含むなら
is_causal=Falseを渡す。未来の token はキャッシュに存在しない。 - 複数 token を一度に処理する chunked Decode では、実際の cache position に合わせた右下揃えの causal mask を作る。
PyTorch の is_causal=True は、非正方行列では左上揃えの mask を使います。そのため [T_q=1, T_k>1] のキャッシュ付き Decode にそのまま True を渡すと、キャッシュ全体ではなく最初の key しか見えません。これは mask の位置合わせの問題であり、MHA、GQA、MQA の違いではありません。
2026年7月時点でも、PyTorch は enable_gqa を実験的機能とし、利用できるバックエンドとテンソル種別に制約を設けています。API があるだけで、その GPU、dtype、バージョンに最適なカーネルが選ばれるとは限りません。実際に選ばれたバックエンドを確認し、プロファイルしてください。参照実装は K と V を必ず展開します。ネイティブ API が展開を避けられるかは、実際に選ばれたバックエンドとカーネル次第です。
付属の実行テストでは、MHA、GQA、MQA の float32/float64 の出力と勾配を比較します。さらに、1 token のキャッシュ付き Decode、非正方行列の causal mask、Value だけが Query/Key と異なる特徴次元を持つ場合も確認します。
23.5 MHA checkpoint を GQA に変える
GQA 論文の uptraining は、設定値を一つ変えるだけではありません。
- MHA の K ヘッドと V ヘッドを group ごとに平均し、小さい GQA projection の初期値にする。Q と output projection は変えない。
- 共有された K/V 表現にモデルを適応させるため、pretraining を続ける。
論文の「5%」は**元の pretraining compute の 5%**であり、「どのモデルでも元データの 5% で fine-tuning すればよい」という意味ではありません。論文は T5.1.1 の実験で、uptrained GQA が MHA に近い品質と MQA に近い速度を得たと報告しました。これは範囲の定まった実験結果であり、すべてのモデルへの保証ではありません。
PyTorch の nn.Linear.weight は [out_features, in_features] です。したがって K または V projection を平均するとき、ヘッド軸は出力側にあります。
def mean_pool_kv_projection(weight, num_query_heads, num_kv_heads):
"""Pool [H_q * D_h, D_model] into [H_kv * D_h, D_model]."""
if weight.ndim != 2:
raise ValueError("weight must be two-dimensional")
if num_query_heads <= 0 or num_kv_heads <= 0:
raise ValueError("head counts must be positive")
if num_query_heads % num_kv_heads != 0:
raise ValueError("num_query_heads must be divisible by num_kv_heads")
if weight.size(0) % num_query_heads != 0:
raise ValueError("weight rows must be divisible by num_query_heads")
head_dim = weight.size(0) // num_query_heads
group_size = num_query_heads // num_kv_heads
return (
weight.reshape(num_query_heads, head_dim, weight.size(1))
.reshape(num_kv_heads, group_size, head_dim, weight.size(1))
.mean(dim=1)
.reshape(num_kv_heads * head_dim, weight.size(1))
)
Linear layer に bias がある場合は、それも同じヘッド group で平均します。変換後は loss、downstream task、long-context quality、serving performance を改めて評価してください。Mean pooling は初期化であり、品質が保たれた証拠ではありません。
23.6 実際の model config はどう違うか
次の表は、論文または公式 config で確認できる少数の snapshot に絞っています。同じ model family でも size や revision によって異なるため、最終的には実際に読み込む checkpoint の config.json を確認します。
| Model snapshot | Q heads | KV heads | 方式 | 1 group の Q heads |
|---|---|---|---|---|
| Llama 2 7B | 32 | 32 | MHA | 1 |
| Llama 2 70B | 64 | 8 | GQA | 8 |
| Mistral-7B-v0.1 | 32 | 8 | GQA | 4 |
| Qwen2-7B | 28 | 4 | GQA | 7 |
Hugging Face 形式の config では、一般に次の field が見つかります。
{
"num_attention_heads": 32,
"num_key_value_heads": 8
}
読み方は機械的です。
- 二つが等しければ MHA。
num_key_value_headsが 1 なら MQA。- その間なら GQA。
- さらに Query ヘッド数が KV ヘッド数で割り切れることを確認する。
「KV 8 ヘッドが sweet spot」という法則はありません。8 は一部の 2/4/8-way tensor parallel 構成できれいに分割できますが、KV ヘッド数が device 数より少ないと、KV ヘッドを複製したり別の sharding を選んだりする実装もあります。品質曲線も model、data、training budget によって変わります。候補となる を決めたら、validation quality、TPOT/throughput、peak memory、目的の並列構成を同時に測ります。
23.7 論文の主張を、その実験範囲で読む
2019年の MQA 論文は、評価した翻訳 task において、品質低下を小さく抑えながら decode を大幅に高速化したと報告しました。2023年の GQA 論文は、T5.1.1 の uptraining 実験において、GQA が MHA に近い品質と MQA に近い速度を得られると報告しました。
どちらにも「評価した設定では」という条件があります。次のことまでは証明していません。
- どの model でも MHA の品質が最高である。
- frontier scale では MQA が必ず受け入れられない。
- GQA は無損失である。
- KV ヘッドを 32 から 8 に減らせば decode が正確に 4 倍速くなる。
- どの architecture にも KV 8 ヘッドが合う。
Architecture が変えるのは、選べる trade-off です。実際の着地点は training、kernel、hardware、batch size、context length、serving framework が一緒に決めます。
23.8 章のまとめ
- KV キャッシュは一般にレイヤーごとに二つの論理テンソルです。 はその軸であり、「1 ヘッドにつき 1 tensor」ではありません。
- MHA、GQA、MQA は の一族で、両端は と です。
- Cache bytes は に比例しますが、実際の同時接続数と速度はその比率だけでは決まりません。
- Native GQA kernel は Query ヘッドを KV ヘッドへ直接対応させます。明示的な
repeat_kvは主に参照実装です。 - MHA checkpoint の mean pooling は GQA の初期化にすぎず、論文ではさらに元の pretraining compute の 5% を uptraining に使いました。
- 時代を超えた「8 ヘッドの default」はありません。正確な config を読み、目的の環境で品質と性能を測ります。
参考文献
- Attention Is All You Need(Vaswani et al., 2017)
- Fast Transformer Decoding: One Write-Head is All You Need(Shazeer, 2019)
- GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints(Ainslie et al., 2023)
- Llama 2: Open Foundation and Fine-Tuned Chat Models(Touvron et al., 2023)
- Mistral-7B-v0.1 公式 config
- Qwen2-7B 公式 config
- PyTorch
scaled_dot_product_attentiondocumentation
次の章へ
GQA が減らすのは、各 token が保存する K/V の量です。しかし、Query がどの過去 token を見るかは変わりません。第24章では次の問いに進みます。完全な履歴を見るのをやめ、選ばれた位置だけに attention すればどうなるのか。これが Sparse Attention の出発点です。