Открыть сервис

Grouped-Query Attention

Grouped-Query Attention (GQA) — это архитектурная оптимизация механизма внимания (attention) в нейронных сетях, в первую очередь в трансформерах, которая уменьшает вычислительную сложность и объём памяти, требуемый для кэширования ключей и значений (KV-кэш) при авторегрессивной генерации текста. GQA является промежуточным решением между Multi-Head Attention (MHA) и Multi-Query Attention (MQA), обеспечивая лучшее соотношение качества и производительности, чем обе крайности.

История и предпосылки

Проблема KV-кэша в авторегрессивных моделях

В трансформерах, используемых для генерации текста (например, GPT, LLaMA), каждый шаг декодирования требует вычисления внимания ко всем предыдущим токенам последовательности. Для ускорения этого процесса модель сохраняет в памяти вычисленные ранее ключи (K) и значения (V) для каждого слоя и каждой головы внимания — это называется KV-кэш. Размер этого кэша линейно растёт с длиной последовательности и числом голов внимания, что становится узким местом при развёртывании больших языковых моделей (LLM) на устройствах с ограниченной памятью.

Multi-Head Attention (MHA)

В стандартном MHA, предложенном в оригинальной работе «Attention Is All You Need» (Vaswani et al., 2017), каждый «голова» внимания имеет свой собственный набор проекций для запросов (Q), ключей (K) и значений (V). Если в модели h голов, то для каждого токена хранится h пар K и V. Это даёт высокую гибкость, но требует больших объёмов памяти для KV-кэша.

Multi-Query Attention (MQA)

Для снижения нагрузки на память была предложена Multi-Query Attention (Shazeer, 2019). В MQA все головы внимания используют один общий набор ключей и значений (одна пара K и V на слой), в то время как запросы (Q) остаются уникальными для каждой головы. Это резко сокращает размер KV-кэша (в h раз), но может приводить к ухудшению качества модели, так как все головы вынуждены извлекать информацию из одного и того же представления.

Появление GQA

Grouped-Query Attention была предложена в 2023 году в статье «GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints» (Ainslie et al., Google Research). Основная идея заключалась в том, чтобы найти компромисс: не одна общая пара K-V (как в MQA) и не h отдельных пар (как в MHA), а группа голов, разделяющих один набор K-V. Это позволяет сохранить часть выразительной способности MHA, одновременно значительно уменьшив размер KV-кэша.

Устройство и принцип работы

Основная концепция

В GQA головы внимания делятся на G групп. Внутри каждой группы все головы используют один и тот же набор ключей и значений. Количество голов в группе обычно обозначается как h / G. Если G = 1, GQA вырождается в MQA. Если G = h, GQA эквивалентна MHA.

Формальное описание

Пусть модель имеет:

  • h — общее число голов внимания.
  • G — число групп (где G делит h нацело).
  • d_k — размерность ключей и значений на одну голову.

В GQA:

  1. Запросы (Q): Каждая из h голов проецирует входное представление в свой собственный запрос. Таким образом, имеется h различных Q-проекций.
  2. Ключи (K) и Значения (V): Создаётся только G различных проекций для ключей и G для значений. Каждая из этих G проекций используется группой из h / G голов.
  3. Вычисление внимания: Для каждой головы i (принадлежащей группе g) вычисляется Attention(Q_i, K_g, V_g).

Сравнение с MHA и MQA

ПараметрMHA (Multi-Head)GQA (Grouped-Query)MQA (Multi-Query)
Число проекций K/Vh (по числу голов)G (по числу групп)1 (одна на слой)
Размер KV-кэшаh n_layers seq_len * d_kG n_layers seq_len * d_k1 n_layers seq_len * d_k
Выразительная способностьВысокая (каждая голова независима)Средняя (головы в группе связаны)Низкая (все головы связаны)
Скорость инференсаМедленнее (больше данных для загрузки)Быстрее (меньше данных для загрузки)Самая быстрая

Типичные конфигурации

На практике часто выбирают G как степень двойки, чтобы упростить аппаратную реализацию. Например:

  • G=2: Две группы голов. Промежуточный вариант.
  • G=4: Четыре группы. Часто используется в моделях семейства LLaMA (например, LLaMA 2 70B использует 8 групп при 64 головах, то есть 8 голов на группу).
  • G=8: Восемь групп.

Преимущества и недостатки

Преимущества

  1. Снижение потребления памяти: Основное преимущество. KV-кэш уменьшается в h / G раз по сравнению с MHA. Это критически важно для развёртывания LLM на GPU с ограниченной видеопамятью (например, потребительские видеокарты) и для обработки длинных контекстов.
  2. Ускорение инференса: Меньший объём данных, который нужно загрузить из памяти в вычислительные ядра (memory-bound operation), приводит к значительному ускорению генерации, особенно при больших размерах батча и длинных последовательностях.
  3. Сохранение качества: GQA обычно показывает результаты, близкие к MHA, и значительно превосходит MQA, особенно в задачах, требующих тонкого понимания контекста (например, суммирование, ответы на вопросы по длинным документам).

Недостатки

  1. Незначительное снижение качества: В некоторых задачах GQA может немного уступать полноценному MHA, так как головы внутри одной группы не могут независимо выбирать разные представления из K и V.
  2. Сложность обучения с нуля: GQA требует либо специальной процедуры обучения (как в оригинальной статье — дообучение из чекпоинта MHA), либо обучения с нуля с архитектурой GQA, что может быть менее стабильно, чем обучение MHA.
  3. Выбор числа групп: Оптимальное число групп G зависит от модели и задачи. Слишком маленькое G (близкое к MQA) может ухудшить качество, слишком большое (близкое к MHA) — не даст существенного выигрыша в памяти.

Применение

GQA стала стандартной архитектурой для многих современных больших языковых моделей, где эффективность инференса является приоритетом.

  • LLaMA 2 и LLaMA 3 (Meta): Модели LLaMA 2 (особенно версии 70B) и LLaMA 3 используют GQA. Например, LLaMA 2 70B имеет 64 головы и 8 групп (G=8), что даёт 8 голов на группу.
  • Gemma (Google): Модели семейства Gemma также применяют GQA.
  • Mistral и Mixtral (Mistral AI): Модели Mistral 7B (G=8 при 32 головах) и Mixtral 8x7B используют GQA.
  • Yi (01.AI): Модели Yi-34B и другие применяют GQA.
  • Falcon (Technology Innovation Institute): Модели Falcon-40B и Falcon-180B используют MQA, который является частным случаем GQA (G=1).

Интересные факты

  • GQA была разработана как метод, позволяющий дообучить уже существующую модель с MHA до архитектуры с GQA без потери качества. Авторы оригинальной статьи показали, что можно взять чекпоинт обученной MHA-модели, усреднить проекции K и V внутри групп (upscaling) и затем дообучить модель с GQA. Это значительно дешевле, чем обучение с нуля.
  • Выбор GQA часто описывают как «лучшее из двух миров» (best of both worlds) — сочетание качества MHA и скорости MQA.
  • В контексте аппаратного ускорения GQA хорошо ложится на архитектуру современных GPU, где чтение из памяти является узким местом. Уменьшение объёма KV-кэша позволяет эффективнее использовать пропускную способность памяти (memory bandwidth).

Критика и ограничения

  • Не универсальность: GQA не является панацеей. Для моделей, работающих с короткими последовательностями (например, классификация текста), выигрыш от GQA может быть незначительным, а усложнение архитектуры — излишним.
  • Компромисс качества: Хотя GQA близка к MHA по качеству, в задачах, требующих максимальной точности (например, научные исследования, юридический анализ), предпочтение может отдаваться MHA, если ресурсы памяти не являются ограничением.
  • Альтернативные подходы: Существуют и другие методы оптимизации KV-кэша, такие как KV-кэш сжатие (например, H2O, StreamingLLM) или использование скользящего окна внимания (sliding window attention). GQA не исключает их, а может использоваться совместно для достижения ещё большей эффективности.

Источники

  • Vaswani, A., et al. (2017). Attention Is All You Need.
  • Shazeer, N. (2019). Fast Transformer Decoding: One Write-Head is All You Need.
  • Ainslie, J., et al. (2023). GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints.
  • Touvron, H., et al. (2023). LLaMA 2: Open Foundation and Fine-Tuned Chat Models.
  • Jiang, A. Q., et al. (2023). Mistral 7B.

BFOmetr — база данных и аналитика по компаниям России.

На главную BFOmetr →