Математика глубокого обучения

KV-кэш

Почему генерация ограничена памятью, а не арифметикой — и что из этого следует

Шаг 87 из 117 · ~28 мин

Наблюдение, из которого всё следует

При авторегрессионной генерации на шаге n+1n+1 ключи и значения для позиций 1n1 \dots n уже были посчитаны на предыдущих шагах, и они не изменились: kik_i зависит только от токена ii. Значит их можно сохранить.

Attn(qn+1,K1:nkn+1,V1:nvn+1)\text{Attn}(q_{n+1}, \htmlData{k=cached}{K_{1:n}} \oplus \htmlData{k=fresh}{k_{n+1}}, \htmlData{k=cached}{V_{1:n}} \oplus \htmlData{k=fresh}{v_{n+1}})

Экономия арифметики огромна: без проекции KK и VV пересчитываются для всех nn позиций на каждом шаге, с кэшем — для одной. Отношение растёт линейно по nn и при n=4096n = 4096 составляет примерно ×1366\times 1366.

Но за это платится памятью, и вот тут начинается интересное.

Сколько это байт

размер кэша=2nLdkvбайт\text{размер кэша} = 2 \cdot n \cdot L \cdot d_{\text{kv}} \cdot \text{байт}

Двойка — за KK и VV; LL — число слоёв; dkv=gdd_{\text{kv}} = g \cdot d — суммарная ширина ключей. Ни одного множителя лишнего, и весь смысл в том, что nn входит линейно, а веса от nn не зависят вовсе.

Отсюда точка пересечения. Веса стека — примерно 12LD212LD^2 (четыре квадратные проекции внимания плюс MLP с четырёхкратным расширением), значит

2nLdkv=12LD2n=6D2dkv=6Dhg2nLd_{\text{kv}} = 12LD^2 \quad\Longrightarrow\quad n^\ast = \frac{6D^2}{d_{\text{kv}}} = 6D\cdot\frac{h}{g}

Потяните длину контекста и посмотрите, когда синяя полоса перерастает серую.

кэш · 2048 МБвеса · 12288 МБ
длина контекста n 4096
групп ключей g 32
кэш
2048 МБ
веса
12288 МБ
сравняются при
24576 токенов
выигрыш против пересчёта
×1366
операций на элемент кэша
1 на элемент
При полном MHA кэш занимает 0.5 МБ на токен: при 4096 токенах это 2 ГБ рядом с 12 ГБ весов. Сдвиньте g к восьми — кэш падает вчетверо, а операций на элемент становится четыре вместо одной.

Каждая голова со своими ключами: операций на прочитанный элемент кэша ровно одна. Это худший возможный случай для арифметической плотности.

Числа для D=4096D = 4096, L=32L = 32, fp16:

конфигурацияна токенпри n=4096n = 4096сравняются с весами
MHA, g=32g = 320.5000.500 МБ2.02.0 ГБ2457624\,576
GQA, g=8g = 80.1250.125 МБ0.50.5 ГБ9830498\,304
MQA, g=1g = 10.0160.016 МБ0.060.06 ГБ786432786\,432

Полмегабайта за токен — это та цифра, из-за которой длинный контекст дорог не по вычислениям, а по памяти на каждого пользователя одновременно.

Главное число: одна операция на элемент

Теперь самое важное, и это не про размер. Посчитаем арифметическую плотность внимания на шаге декодирования: сколько умножений-сложений приходится на один прочитанный из памяти элемент.

За шаг внимание в одном слое читает 2ndkv2 n d_{\text{kv}} чисел кэша и делает 2nD2nD умножений-сложений. Отношение:

2nD2ndkv=Ddkv=hg\frac{2nD}{2nd_{\text{kv}}} = \frac{D}{d_{\text{kv}}} = \frac{h}{g}

При обычном MHA (g=hg = h) это ровно единица. Одно умножение на одно прочитанное число.

Почему это приговор: у современных ускорителей отношение арифметической производительности к пропускной способности памяти — порядка ста–трёхсот операций на прочитанный байт. Значит при плотности 11 вычислитель простаивает, ожидая память, и ускорение арифметики не даёт ничего. Генерация ограничена чтением кэша.

схемаh/gh/gплотность
MHA1111
GQA, g=h/4g = h/44444
GQA, g=h/8g = h/88888
MQAhh3232 при h=32h = 32

Вот и полное объяснение, зачем существуют MQA и GQA — то, что урок 080 оставил без причины. Они не ускоряют арифметику (её столько же) и не уменьшают число голов (запросов по-прежнему hh). Они делят объём читаемой памяти на h/gh/g, и ровно во столько же раз поднимают плотность. Экономия кэша — приятное следствие, а причина в плотности.

Почему обучение — другая задача

Стоит сопоставить, потому что интуиция из обучения здесь подводит.

обучениегенерация
токенов за проходтысячиодин
матричные операцииматрица × матрицаматрица × вектор
арифметическая плотностьвысокаяоколо единицы
что ограничиваетарифметикапамять
помогает FlashAttentionдамало
помогает GQAпочти нетсильно

Причина в третьей строке. При обучении веса читаются один раз на весь батч, поэтому их чтение амортизируется; при генерации токен один, и каждое число весов и кэша используется ровно один раз. Умножение матрицы на вектор ограничено памятью по природе, а не из-за реализации.

Отсюда и батчирование запросов при сервинге: несколько независимых последовательностей одновременно возвращают плотность, потому что веса читаются один раз на всю группу. Кэш при этом не амортизируется — он у каждой последовательности свой, и именно он ограничивает, сколько запросов уместится.

Что ещё делают с кэшем

приёмидеяцена
GQA / MQAгруппы вместо головнемного выразительности
квантование кэша в int8вдвое меньше байтнебольшая потеря качества
скользящее окнохранить только последние wwзабывание дальше ww
PagedAttentionстраницы вместо непрерывных буферовсложность реализации
выгрузка в CPUпамять дешевлепропускная способность падает

Первые четыре снижают либо объём чтения, либо фрагментацию — то есть работают против того же узкого места. Последний пункт — противоположная ставка, и она проигрывает как раз потому, что проблема в пропускной способности, а не в ёмкости.

Итог

  • Кэш превращает O(n)O(n) пересчётов на шаг в O(1)O(1) — выигрыш растёт линейно по контексту.
  • Он растёт линейно по nn, а веса не растут вовсе, поэтому существует вычислимая точка n=6Dh/gn^\ast = 6D\,h/g, где кэш становится больше модели.
  • Арифметическая плотность внимания при декодировании равна h/gh/g, то есть единице для MHA. Это и есть причина существования GQA и MQA.
  • Генерация ограничена памятью, обучение — арифметикой. Приёмы для одного почти не помогают другому.

Источники

Проверки

0 из 2
  1. Что ограничивает генерацию

    Отметьте все верные утверждения о KV-кэше и декодировании.

  2. Бюджет кэша и плотность

    Реализуйте cache_facts(layers, heads, head_dim, groups, context, bytes_per_value). Обозначим D=hdD = h \cdot d (ширина модели) и dkv=gdd_{\text{kv}} = g \cdot d (суммарная ширина ключей). Верните [cache_megabytes, crossover, intensity, speedup]:

    • cache_megabytes = 2nLdkvбайт2 n L\, d_{\text{kv}} \cdot \text{байт}, делённое на 102421024^2. Двойка — за ключи и значения;
    • crossover = 6D2dkv\dfrac{6D^2}{d_{\text{kv}}} — длина контекста, при которой кэш сравняется с весами стека, оценёнными как 12LD212LD^2 элементов;
    • intensity = D/dkvD / d_{\text{kv}} — умножений-сложений на один прочитанный из кэша элемент;
    • speedup — отношение работы без кэша к работе с кэшем на один новый токен в одном слое:

    2D2+2nDdkv+2nD2D2+2Ddkv+2nD\frac{2D^2 + 2n D d_{\text{kv}} + 2nD}{2D^2 + 2D d_{\text{kv}} + 2nD}

    Числитель отличается от знаменателя одним множителем nn во втором слагаемом: без кэша ключи и значения пересчитываются для всех nn позиций, с кэшем — для одной.

    Проверить себя можно так: intensity обязана равняться h/gh/g, и при g=hg = h это ровно единица.

    функция cache_facts

    Загрузка редактора…

    Ctrl/⌘ + Enter