56 GQA / MQA
Контекст
При обслуживании LLM основная память уходит на KV-кэш (см. #49), и он пропорционален числу голов, для которых мы храним ключи и значения. Полное multi-head attention (#32) держит отдельные K, V для каждой из h голов — дорого по памяти на длинном контексте.
Идея и механизм
Наблюдение: query-голов нужно много (они задают разные «вопросы»), а вот ключи и значения можно делить. MQA доводит это до предела: одна общая K/V-голова на все query — кэш падает в h раз, но модель теряет выразительность и хуже/нестабильнее обучается. GQA — компромисс: делим h query-голов на g групп, каждая группа делит одну K/V-голову. При g = 1 это MQA, при g = h — обычное MHA. Бонус: GQA можно дёшево «доучить» (uptrain) из готового MHA-чекпойнта, а не тренировать с нуля.
линейная алгебра · системы Откуда экономия KV-кэша
Размер KV-кэша на запрос пропорционален числу KV-голов \( n_{kv} \) (ключи и значения на каждый слой, KV-голову и токен):
Меняя только \( n_{kv} \), получаем весь спектр — от полного MHA до предельного MQA:
Например, 32 query-головы и 8 KV-групп → каждая четвёрка query делит одну K/V → 4× меньше KV-кэша при качестве почти как у MHA. Число query-голов (а значит, стоимость самого attention-матмула) не меняется — экономия именно в памяти кэша, то есть в том, сколько запросов влезет в батч.
PyTorch GQA: «размножаем» g KV-голов до h для матмула
import torch
# q: [B, h, T, d]; k,v: [B, g, T, d] (g KV-групп, g делит h)
def gqa(q, k, v, h, g):
rep = h // g # сколько query на одну KV-голову
k = k.repeat_interleave(rep, dim=1) # [B, h, T, d] — общий K на группу
v = v.repeat_interleave(rep, dim=1)
a = (q @ k.transpose(-1,-2)) / q.size(-1)**0.5
return a.softmax(-1) @ v # храним только g KV-голов, считаем как h
Почему это важно
GQA сделала длинный контекст дешёвым по памяти и стала де-факто стандартом почти всех современных open-LLM (LLaMA-2 70B, LLaMA-3, Mistral). Это часть той же борьбы за KV-кэш, что vLLM (#49), prefix caching (#54) и MLA (#53) — только здесь кэш ужимают уменьшая число KV-голов.
Связи
GQA/MQA — прямая модификация multi-head attention из #32: та же механика, но K/V-головы делятся между query. Меняется не идея внимания, а её «бухгалтерия» памяти на инференсе.
Обе атакуют стоимость внимания, но с разных сторон: FlashAttention сокращает трафик памяти при вычислении attention, GQA — размер KV-кэша между шагами. В проде их комбинируют.
MLA из DeepSeek — другой способ ужать тот же кэш: GQA сокращает число KV-голов (делит их), а MLA сжимает каждую KV в низкоранговый латент. Разные оси уменьшения одного боттлнека.
Вопросы пытливого ума
Почему делить K/V почти не роняет качество, а делить query — роняет?
Query-головы задают разные вопросы к контексту — их разнообразие несёт много информации, и урезать его дорого. А ключи и значения — это общее «содержимое», в котором эмпирически больше избыточности: несколько query вполне могут смотреть в один и тот же K/V без большой потери. Поэтому сокращают именно KV-головы.
Если MQA (g=1) так компактен, зачем вообще GQA?
MQA слишком агрессивен: одна K/V-голова на всех теряет выразительность, даёт заметную просадку качества и нестабильно обучается (особенно при up-training больших моделей). GQA оставляет несколько групп — этого хватает, чтобы удержать качество почти на уровне MHA, но кэш всё ещё в разы меньше. g — ручка размена «память ↔ качество».
Можно ли получить GQA-модель, не обучая с нуля?
Да — это ключевой практический вклад статьи. Готовый MHA-чекпойнт «конвертируют»: усредняют K/V-головы внутри каждой группы в одну, а затем uptrain — коротко дообучают (малая доля исходного бюджета). Так почти любую существующую MHA-модель дёшево переводят на GQA, не платя за полное обучение.
Что читать в оригинале
Читать выборочно: из GQA (arXiv 2305.13245) — идею групп и uptraining из MHA-чекпойнта, кривую «качество vs число групп»; из MQA (Shazeer, 2019) — исходный тезис «одной write-head достаточно» и почему кэш был узким местом ещё тогда.