Эпоха 6 · Генеративка и системы · 2019 / 2023

56 GQA / MQA

Fast Transformer Decoding: One Write-Head is All You Need (MQA · Shazeer, 2019) · GQA: Training Generalized Multi-Query Transformer Models… (Ainslie и др. · Google, 2023)
🟧 оригинал выборочно~1 чоригинал ↗
Суть за 20 секунд. KV-кэш растёт с числом голов внимания — а это узкое место инференса. MQA (2019): все query-головы делят ОДНУ K/V-голову → крошечный кэш, но просадка качества и нестабильное обучение. GQA (2023): золотая середина — h query-голов разбиты на g групп, в группе общая K/V-голова. Тюнингуемый g (1 = MQA, h = MHA), обычно ~4× меньше KV при качестве близком к полному MHA. Стандарт LLaMA-2/3, Mistral.

Контекст

При обслуживании 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-голову и токен):

\[ \text{KV cache} \;\propto\; 2\,L\,n_{kv}\,d_{h}\,T \]

Меняя только \( n_{kv} \), получаем весь спектр — от полного MHA до предельного MQA:

\[ \text{MHA: } n_{kv}=h \qquad \text{GQA: } n_{kv}=g \qquad \text{MQA: } n_{kv}=1 \quad\Rightarrow\quad \text{cache} \downarrow\ \tfrac{h}{g}\times \]

Например, 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
MHA (n_kv = h) QKV GQA (n_kv = g) ↑2 группы MQA (n_kv = 1) 1 на всех меньше KV-голов → меньше кэш → больше запросов в батч query-головы (синие) не меняются; делятся только K/V (число снизу)
Число query-голов постоянно; уменьшается число KV-голов, которые они делят: MHA — по одной на каждую, GQA — по одной на группу, MQA — одна на всех. Меньше KV-голов = меньше кэша.
Аналогия. Переводчики (query-головы) и справочники (K/V). В MHA у каждого переводчика свой полный справочник — точно, но громоздко. В MQA один общий справочник на всех — компактно, но у стойки толпа и детали теряются. GQA даёт по справочнику на небольшую группу переводчиков: почти так же точно, но справочников в разы меньше — их и «возить» (держать в памяти) дешевле.

Почему это важно

GQA сделала длинный контекст дешёвым по памяти и стала де-факто стандартом почти всех современных open-LLM (LLaMA-2 70B, LLaMA-3, Mistral). Это часть той же борьбы за KV-кэш, что vLLM (#49), prefix caching (#54) и MLA (#53) — только здесь кэш ужимают уменьшая число KV-голов.

Связи

← упрощает32. Transformer

GQA/MQA — прямая модификация multi-head attention из #32: та же механика, но K/V-головы делятся между query. Меняется не идея внимания, а её «бухгалтерия» памяти на инференсе.

↔ другая ось48. FlashAttention

Обе атакуют стоимость внимания, но с разных сторон: FlashAttention сокращает трафик памяти при вычислении attention, GQA — размер KV-кэша между шагами. В проде их комбинируют.

↔ ещё сильнее жмёт53. DeepSeek V3 / R1

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 достаточно» и почему кэш был узким местом ещё тогда.