Skip to content

attention変種(MQA/GQA/MLA)

実装: llm/gqa/ / 実行: go test ./llm/gqa/

マルチヘッド attention はヘッドごとに独立の K/V を持つが、生成時にキャッシュする K/V のメモリはヘッド数に比例して膨らみ、長文生成のボトルネックになる。K/V だけをヘッド間で共有するのが MQA と GQA だ。この章では 3 方式を 1 つの実装で書き、違いが「Q ヘッドがどの K/V を引くか」の対応表だけであることを確かめる。圧縮で削る MLA も最後に見る。

この章で作るもの

mini-GPT で組んだマルチヘッド attention(MHA)は、ヘッドごとに独立の Q/K/V を持っていた。表現力の点ではそれでよいが、実際に文章を生成させると別の問題が現れる。生成は 1 トークンずつ進み、各ステップで過去全トークンの K と V が要る。これを毎回再計算しないよう保存しておくのが KV キャッシュで(あとの推論高速化の章で実装する)、そのサイズは次の式で決まる:

KVキャッシュ = 2 (KとV) × 系列長 × KVヘッド数 × ヘッド次元 × 層数 × バイト/値

MHA では KV ヘッド数 = ヘッド数なので、32 ヘッド × 4096 次元 × 32 層のモデルで 8k トークンを保持すると、1 リクエストだけで GB 級のメモリを食う。同時に多数のリクエストをさばく推論サーバでは、GPU メモリの大半がモデル本体ではなくこのキャッシュに消える。

この式でモデル側から動かせるのは KV ヘッド数だ。Q ヘッドは 32 のまま、K/V だけを共有して減らす。全ヘッドで 1 組まで減らすのが MQA(multi-query attention)、グループごとに 1 組持つ折衷が GQA(grouped-query attention)になる。

MHA (32Q : 32KV)   Q0→KV0  Q1→KV1  Q2→KV2  Q3→KV3   キャッシュ 1
GQA (32Q :  8KV)   Q0→KV0  Q1→KV0  Q2→KV1  Q3→KV1   キャッシュ 1/4
MQA (32Q :  1KV)   Q0→KV0  Q1→KV0  Q2→KV0  Q3→KV0   キャッシュ 1/32
                   (図は 4 ヘッドに縮めた模式)
3 方式の違いは Q ヘッドと K/V の対応表だけ。MHA は 1 対 1、MQA は全 Q が 1 組を共有、GQA はグループ単位で共有する。KV キャッシュは K/V の組数に比例して縮む

順に見ていく。

  1. 3 方式は同じ計算: 違いは「Q ヘッド h がどの KV ヘッドを引くか」の対応表だけ。attention の式そのものは変わらない
  2. キャッシュの式に NHeads は現れない: 効くのは KV ヘッド数。そこだけ減らせばキャッシュが線形に縮む
  3. 品質と削減の折衷が GQA: MQA は最大削減だが品質が下がりやすい。グループ共有の GQA が現在の主流

① 構成: KV ヘッド数という設計変数

実装は 3 方式を別々に書かず、KV ヘッド数を変数にした一般形として書く。NKVHeads = NHeads で MHA、= 1 で MQA、その間が GQA だ:

go

// Config は attention のヘッド構成。
// NKVHeads = NHeads で MHA、= 1 で MQA、その間が GQA になる。
type Config struct {
	DModel   int // 埋め込み次元
	NHeads   int // Q ヘッド数
	NKVHeads int // K/V ヘッド数(NHeads の約数)
}

// KVCacheFloats は系列長 seqLen まで生成したときにキャッシュする float 数。
// K と V の 2 本 × 系列長 × KV ヘッド数 × ヘッド次元。NHeads は現れない。
// これが「K/V を共有するとキャッシュが縮む」ことの式そのもの。
func (c Config) KVCacheFloats(seqLen int) int {
	return 2 * seqLen * c.NKVHeads * (c.DModel / c.NHeads)
}

KVCacheFloats が上の式の実装で、NHeads がどこにも現れないことがそのまま「Q ヘッドを減らさずキャッシュだけ減らせる」理由になっている。テストでは 32:8 の GQA が MHA の 1/4、MQA が 1/32 になることを固定した。

② forward: 共有の実体は対応表

計算本体を見ると、K/V の射影は KV ヘッド数ぶんしか行われない。各 Q ヘッドは KVHeadFor で自分の共有先を引く:

go

// KVHeadFor は Q ヘッド h が共有する KV ヘッドの番号を返す。
// MHA なら h 自身、MQA なら常に 0、GQA ならグループ番号。
func (a *Attention) KVHeadFor(h int) int {
	return h / (a.cfg.NHeads / a.cfg.NKVHeads)
}

// Forward は causal self-attention を計算し (seq, DModel) を返す。
// K/V は KV ヘッド数ぶんしか作らず、各 Q ヘッドは対応表で共有先を引く。
func (a *Attention) Forward(x *tensor.Tensor) *tensor.Tensor {
	// K/V は KV ヘッドごとに 1 回だけ射影する(ここが共有の実体)。
	ks := make([]*tensor.Tensor, a.cfg.NKVHeads)
	vs := make([]*tensor.Tensor, a.cfg.NKVHeads)
	for h := 0; h < a.cfg.NKVHeads; h++ {
		ks[h] = tensor.MatMul(x, a.wk[h])
		vs[h] = tensor.MatMul(x, a.wv[h])
	}

	out := tensor.New(x.Rows, a.cfg.DModel)
	for h := 0; h < a.cfg.NHeads; h++ {
		q := tensor.MatMul(x, a.wq[h])
		kv := a.KVHeadFor(h)
		head := causalAttend(q, ks[kv], vs[kv], a.dHead)
		// ヘッド出力を (seq, DModel) の担当区画に連結する。
		for r := 0; r < head.Rows; r++ {
			for c := 0; c < head.Cols; c++ {
				out.Set(r, h*a.dHead+c, head.At(r, c))
			}
		}
	}
	return out
}

// causalAttend は softmax(Q·Kᵀ/√d + 因果マスク)·V。attention 編と同じ計算。
func causalAttend(q, k, v *tensor.Tensor, dHead int) *tensor.Tensor {
	scores := tensor.MatMul(q, transpose(k))
	scale := float32(1.0 / math.Sqrt(float64(dHead)))
	negInf := float32(math.Inf(-1))
	for i := 0; i < scores.Rows; i++ {
		for j := 0; j < scores.Cols; j++ {
			if j > i {
				scores.Set(i, j, negInf)
			} else {
				scores.Set(i, j, scores.At(i, j)*scale)
			}
		}
	}
	return tensor.MatMul(tensor.SoftmaxRows(scores), v)
}

MHA との差分はこれだけだ。このことはテストでも確かめている。全 KV ヘッドの重みを同一にすると、MHA・GQA・MQA の出力は完全に一致する。3 方式の違いは K/V の中身が何通りあるかだけで、attention の計算そのものはどれも同じということだ。

品質面の直観も同じ場所から出る。ヘッドの多様性のうち「何を探すか」(Q)は全ヘッドぶん残り、「何を持っているか」(K/V)の多様性だけが減る。MQA まで削ると品質低下が測定できるレベルで現れることがあり、Llama 2 70B 以降の主要モデルは 4〜8 ヘッドに 1 組の GQA に落ち着いている。

③ MLA: 共有ではなく圧縮する

DeepSeek-V2/V3 はさらに別の路線を取った。MLA(multi-head latent attention)は K/V をヘッド間で共有するのではなく、低ランクの潜在ベクトルに圧縮してキャッシュする。K/V を作る前の中間表現(数百次元)だけを保存し、attention 時にそこから各ヘッドの K/V を復元する。

  • GQA との違い: GQA は「K/V の種類を減らす」、MLA は「K/V の元を細くする」。MLA はヘッドごとの多様性を保ったままキャッシュを削れる
  • 代償: 復元のための行列積が毎ステップ増える。メモリと引き換えに計算を払う設計で、メモリ律速の推論サーバでは得になる

DeepSeek-V2 の報告では、MLA は MHA 比で KV キャッシュを 93% 削減しつつ品質を上回った。ここでは仕組みの解説に留め、実装は共有系(MQA/GQA)までとした。

動かす

下のデモは「対応表」と「キャッシュの式」をそのまま操作できる。MHA / GQA / MQA を切り替えると、8 つの Q ヘッドから K/V への線のつながりが変わり、系列長を伸ばすと KV キャッシュのバーが方式ごとに違う速さで伸びる。同じ系列長でもキャッシュが 1/4、1/8 になることが見て取れる。

デモattention変種(KVキャッシュ)MHA · 1,024 tok
MHAGQAMQA図は 8Q の模式 / 実測は 32Q・128次元・32層・fp16
Q ヘッドと K/V の対応(8Q : 8KV)
Q0
↓ 共有
KV0
Q1
↓ 共有
KV1
Q2
↓ 共有
KV2
Q3
↓ 共有
KV3
Q4
↓ 共有
KV4
Q5
↓ 共有
KV5
Q6
↓ 共有
KV6
Q7
↓ 共有
KV7
KV キャッシュ(1 リクエストあたり) 系列長 1,024
MHA (32KV)512 MiB
GQA (8KV)128 MiB
MQA (1KV)16 MiB

MHA: 全 Q ヘッドが自分専用の K/V を持つ。表現力は基準だが、キャッシュもヘッド数ぶんまるごと。系列長 1,024 で 512 MiB

1 / 6

キャッシュの式 2 × 系列長 × KVヘッド数 × ヘッド次元 × 層数 に Q ヘッド数は現れない。 だから K/V の共有だけでメモリが 1/4(GQA)や 1/32(MQA)に縮む。 attention の計算そのものは 3 方式で同じで、違いは対応表だけ。

設計の観点

  • なぜ K/V だけ共有するか: 生成時にキャッシュされるのは K と V だけ(Q は毎ステップ新トークンのものを作れば済む)。だから Q を減らしてもメモリは減らず、K/V を減らせば線形に減る
  • 品質への影響の非対称: Q の多様性は「どこを見るか」の多様性で、これを残せば各ヘッドは共有 K/V から違う場所を引ける。K/V の多様性削減の方が影響が小さいという経験則が GQA の根拠
  • GQA のグループ数の選び方: 実務では GPU のテンソル並列数と揃えることが多い(KV ヘッド 8 で 8 GPU に 1 つずつ)。ハードウェア構成が設計に染み出す例
  • MLA との使い分け: 共有(GQA)は実装が軽く既存構造を保つ。圧縮(MLA)は削減率と品質で勝るが、復元計算と実装の複雑さを払う
  • アップトレイン: 既存 MHA モデルの K/V ヘッドを平均して GQA 化し、少量の追加学習で回復させられる(GQA 論文の手法)。ゼロから学習し直さなくてよい

メリット・デメリットと実例

方式KVキャッシュ品質実例
MHA1(基準)基準GPT-2/3、初期 Llama、mini-GPT
MQA1/NHeads低下が出やすいPaLM、Falcon-7B、Gemini 1.0 Pro 系の一部
GQANKVHeads/NHeadsほぼ維持Llama 2 70B / Llama 3、Mistral、Qwen 2 以降
MLA数%まで圧縮維持〜向上の報告DeepSeek-V2 / V3 / R1

裏どり:

  • MQA(Shazeer 2019): "Fast Transformer Decoding" で提案。1 人の著者による 4 ページの論文が、のちの推論効率化路線の起点になった
  • GQA(Google 2023): MQA の品質低下と MHA のメモリの間を取る論文。既存モデルからのアップトレイン手順も示し、Llama 2 70B が採用して標準化した
  • Llama 3(2024): 8B モデルまで GQA(32Q:8KV)を採用。小型でも長コンテキスト運用でキャッシュが効くため
  • DeepSeek-V2(2024): MLA の初出。KV キャッシュ 93% 削減の報告とともに、共有ではなく圧縮という第 3 の路線を示した

簡略化したこと

  • 出力射影(Wo)なし: ヘッド連結までを実装。実物は連結後にもう 1 枚射影が入る
  • KV キャッシュ自体は未実装: この章はキャッシュのサイズを決める構造の話。キャッシュを実際に持って生成を速くする実装は推論高速化の章
  • MLA は解説のみ: 低ランク圧縮と復元の実装は行わない
  • RoPE との結合なし: 実物は Q/K に RoPE を挟む。MLA では RoPE との両立に追加の工夫が要る(decoupled RoPE)

参考資料