Multi-query och grouped-query attention
Kunna förklara hur MQA/GQA minskar KV-cache och vad som offras.
Förkunskaper
- EKV-cachekrävs
Intuition
I vanlig multi-head attention (MHA) har varje huvud egna Q, K och V. KV-cachen skalar med antalet huvuden — och blir minnesboven vid lång kontext.
MQA (multi-query): alla huvuden delar ett K/V-par. Cachen krymper med antalet huvuden (32×), men kvaliteten sjunker märkbart.
GQA (grouped-query): huvudena delas i grupper som delar K/V. Med 32 Q-huvuden och 8 KV-grupper blir cachen 4× mindre med nästan oförändrad kvalitet.
GQA är standard i Llama 3, Mistral och Qwen — en av de få ändringar som ger stor vinst till nästan inget pris.
Formellt
KV-cachens storlek:
där = lager, = antal KV-huvuden (inte Q-huvuden), = sekvenslängd, = batch.
Exempel, 7B-modell (L = 32, d_head = 128, fp16, T = 8 192, B = 1):
| Variant | H_kv | Cache |
|---|---|---|
| MHA | 32 | 4,3 GB |
| GQA (4 grupper) | 8 | 1,1 GB |
| MQA | 1 | 0,13 GB |
Med GQA ryms fyra gånger fler samtidiga användare på samma kort — vilket direkt översätts till kostnad per användare.
Vad som offras: uttrycksförmåga. Varje grupp måste dela på samma nycklar och värden, så huvuden inom en grupp kan inte specialisera sig lika fritt. Empiriskt är GQA med 4–8 grupper nästan omöjlig att skilja från MHA, medan MQA syns i kvaliteten.
Uptraining: en MHA-modell kan konverteras till GQA genom att medelvärdesbilda K/V-projektionerna inom varje grupp och sedan fortsätta träna på ~5 % av det ursprungliga beräkningsbudgeten (Ainslie m.fl. 2023).
Kod
def kv_cache_gb(lager=32, kv_huvuden=32, d_head=128, seq=8192, batch=1, bytes_per_tal=2):
return 2 * lager * kv_huvuden * d_head * seq * batch * bytes_per_tal / 1e9
for namn, h in (("MHA", 32), ("GQA-8", 8), ("GQA-4", 4), ("MQA", 1)):
print(f"{namn:6s} {kv_cache_gb(kv_huvuden=h):6.2f} GB")
# MHA 4.29 GB
# GQA-8 1.07 GB
# GQA-4 0.54 GB
# MQA 0.13 GB
# Hur många samtidiga användare ryms på ett 80 GB-kort efter modellvikterna (14 GB)?
for namn, h in (("MHA", 32), ("GQA-8", 8)):
print(namn, int((80 - 14) / kv_cache_gb(kv_huvuden=h)))
# MHA 15
# GQA-8 61
Behärskning innebär
- Förklarar MHA, MQA och GQA
- Räknar ut KV-cachens storlek för varje variant
- Vet vad som offras
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- arXiv — GQA: Training Generalized Multi-Query Transformer Models — arXiv (öppen åtkomst; licens per artikel)
- arXiv — Fast Transformer Decoding: One Write-Head is All You Need (MQA) — arXiv (öppen åtkomst; licens per artikel)