一文でまとめると: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]

ヘッド数 HKVH_{KV} はテンソルの一つの軸です。保存量を決めるのは、この軸の長さです。

bytes=B×L×2×HKV×S×dh×b\text{bytes}=B\times L\times 2\times H_{KV}\times S\times d_h\times b

ここで BB は batch size、LL はレイヤー数、SS はキャッシュ済み token 数、dhd_h は head dimension、bb は要素あたりの byte 数です。2 は K と V を表します。

Query ヘッド数と KV ヘッド数で定義される MHA、GQA、MQA の一つの族

この三つは、二つの数字だけでまとめられます。

  • HQH_Q:Query ヘッド数
  • HKVH_{KV}:Key ヘッド数および Value ヘッド数

HQH_QHKVH_{KV} で割り切れる必要があります。各 K/V ペアを共有する Query ヘッド数は次の通りです。

R=HQHKVR=\frac{H_Q}{H_{KV}}

したがって:

方式条件Query 8 ヘッドの例
MHAHKV=HQH_{KV}=H_QKV 8 ヘッド、1 group に Q 1 ヘッド
GQA1<HKV<HQ1<H_{KV}<H_QKV 2 ヘッド、1 group に Q 4 ヘッド
MQAHKV=1H_{KV}=1KV 1 ヘッドを Q 8 ヘッドで共有

MHA と MQA は別々の発明ではなく、GQA もその間に無理やり挟んだ第四の仕組みではありません。一つのパラメータ族の両端と、その間の領域です。


23.2 直感:八人の蒔絵師と図案帳

八人の蒔絵師が、修復途中の硯箱を調べているとします。ある人は螺鈿の欠けを見つけたい。別の人は梨地粉の粒度を確かめたい。さらに別の人は蓋の反りを調べたい。

  • MHA:八人がそれぞれ自分の図案帳を持ち、異なる問いを立て、異なる方法で硯箱の特徴を記録します。
  • MQA:問いは八通りのままですが、全員が一冊の共通図案帳を参照します。
  • GQA:職人をいくつかの組に分け、組ごとに一冊の図案帳を共有します。別の組には別の帳面があります。

問い方が Q projection、過去を索引して持ち運ぶ方法が K と V に対応します。K/V を共有すると保存する履歴は減りますが、Query ヘッドがすべて同じになるわけではありません。Q の行と attention 出力は別々のままです。

Grouped-Query Attention における Q、K、V と KV キャッシュの形状

三方式の形状は同じ記法で表せます。

q: [B, H_Q,  T_q, d_h]
k: [B, H_KV, T_k, d_h]
v: [B, H_KV, T_k, d_h]

HQ=8H_Q=8HKV=2H_{KV}=2 なら、対応は次の通りです。

Q0 Q1 Q2 Q3  -> KV0
Q4 Q5 Q6 Q7  -> KV1

概念上は、各 KV ヘッドを RR 回繰り返してから通常の MHA を計算できます。これは説明や参照実装には便利です。ただし効率のよい kernel は、キャッシュを実際に膨らませず、Query ヘッドを対応する KV ヘッドへ直接割り当てます。


23.3 何がどれだけ小さくなるのか

23.3.1 KV キャッシュ

第22章と同じ条件を使います。32 レイヤー、dh=128d_h=128、1,024 tokens、FP16 または BF16(1 要素 2 bytes)、batch size 1 です。

方式HKVH_{KV}正確な byte 数二進単位MHA 比
MHA32536,870,912512 MiB100%
GQA8134,217,728128 MiB25%
MQA116,777,21616 MiB3.125%
MHA、GQA、MQA の KV キャッシュと attention projection パラメータの比較

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 を D=HQdhD=H_Qd_h とします。Q と output projection はそれぞれ D2D^2、K と V はそれぞれ D(HKVdh)D(H_{KV}d_h) 個の weight を持ちます。

Pattn=2D2+2D2HKVHQP_{attn}=2D^2+2D^2\frac{H_{KV}}{H_Q}

HQ=32H_Q=32 の場合:

方式Attention projection weights
MHA、HKV=32H_{KV}=324D24D^2
GQA、HKV=8H_{KV}=82.5D22.5D^2
MQA、HKV=1H_{KV}=12.0625D22.0625D^2

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 は、設定値を一つ変えるだけではありません。

  1. MHA の K ヘッドと V ヘッドを group ごとに平均し、小さい GQA projection の初期値にする。Q と output projection は変えない。
  2. 共有された K/V 表現にモデルを適応させるため、pretraining を続ける。
MHA の K/V ヘッドを平均し、追加 pretraining で GQA に変換する流れ

論文の「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 snapshotQ headsKV heads方式1 group の Q heads
Llama 2 7B3232MHA1
Llama 2 70B648GQA8
Mistral-7B-v0.1328GQA4
Qwen2-7B284GQA7
Model config の Query ヘッド数と KV ヘッド数を読み、native kernel 対応を確認する

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 によって変わります。候補となる HKVH_{KV} を決めたら、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 章のまとめ

  1. KV キャッシュは一般にレイヤーごとに二つの論理テンソルです。HKVH_{KV} はその軸であり、「1 ヘッドにつき 1 tensor」ではありません。
  2. MHA、GQA、MQA は HQ/HKVH_Q/H_{KV} の一族で、両端は HKV=HQH_{KV}=H_QHKV=1H_{KV}=1 です。
  3. Cache bytes は HKVH_{KV} に比例しますが、実際の同時接続数と速度はその比率だけでは決まりません。
  4. Native GQA kernel は Query ヘッドを KV ヘッドへ直接対応させます。明示的な repeat_kv は主に参照実装です。
  5. MHA checkpoint の mean pooling は GQA の初期化にすぎず、論文ではさらに元の pretraining compute の 5% を uptraining に使いました。
  6. 時代を超えた「8 ヘッドの default」はありません。正確な config を読み、目的の環境で品質と性能を測ります。

参考文献

  1. Attention Is All You Need(Vaswani et al., 2017)
  2. Fast Transformer Decoding: One Write-Head is All You Need(Shazeer, 2019)
  3. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints(Ainslie et al., 2023)
  4. Llama 2: Open Foundation and Fine-Tuned Chat Models(Touvron et al., 2023)
  5. Mistral-7B-v0.1 公式 config
  6. Qwen2-7B 公式 config
  7. PyTorch scaled_dot_product_attention documentation

次の章へ

GQA が減らすのは、各 token が保存する K/V の量です。しかし、Query がどの過去 token を見るかは変わりません。第24章では次の問いに進みます。完全な履歴を見るのをやめ、選ばれた位置だけに attention すればどうなるのか。これが Sparse Attention の出発点です。

このページを引用する
Zhang, Wayland (2026). 第23章:MHA から MQA、そして GQA へ. In Transformer アーキテクチャ:直感から実装まで. https://waylandz.com/llm-transformer-book-ja/chapter-23-mha-mqa-gqa/
@incollection{zhang2026transformer_ja_chapter-23-mha-mqa-gqa,
  author = {Zhang, Wayland},
  title = {第23章:MHA から MQA、そして GQA へ},
  booktitle = {Transformer アーキテクチャ:直感から実装まで},
  year = {2026},
  url = {https://waylandz.com/llm-transformer-book-ja/chapter-23-mha-mqa-gqa/}
}