Память attention на практике: матрица оценок, буферы и KV-cache
Автор: Казачкин Даниил Михайлович · Обновлено
После этого практикума вы сможете оценить размер конкретного буфера attention, объяснить рост памяти при увеличении контекста и отличить матрицу оценок от KV-cache. Все расчёты выполняются стандартным Python 3 без GPU. Результат — таблица байтов и маленький замер выделений памяти, а не оценка скорости видеокарты.
Предполагается, что вы знаете роли Q, K и V из [урока о Transformer](/lessons/without-university/generative-ai-foundations/generative-ai-03). Исследовательская идея блочного вычисления и проверка его точности разобраны в [статье о FlashAttention](/research/flashattention-exact-attention-memory-io). Здесь решаем прикладной вопрос: какой именно объект вырос и какой эксперимент проверить первым при нехватке памяти.
Сначала запишите формы массивов
Пусть B — размер batch, H — число голов, N — длина последовательности, D — размер одной головы, S — число байтов на элемент. Для обычного multi-head attention размеры Q, K и V равны B × H × N × D. Одна материализованная матрица оценок имеет размер B × H × N × N. Выход attention снова содержит B × H × N × D элементов.
Разница в последнем множителе определяет рост: удвоение N удваивает Q, K и V, но увеличивает полную матрицу оценок в четыре раза. Это расчёт одного типа буфера, не всей модели. Веса, другие активации, градиенты, оптимизатор и служебные выделения пока исключены.
Калькулятор без больших выделений
Сохраните пример как attention_budget.py и запустите python3 attention_budget.py. Он оперирует целыми числами и не создаёт гигабайтных массивов.
MIB = 1024 ** 2
def buffers(batch, heads, length, head_dim, element_bytes):
vector = batch * heads * length * head_dim * element_bytes
return {
"scores": batch * heads * length * length * element_bytes,
"qkv": 3 * vector,
"output": vector,
}
def kv_cache(batch, layers, kv_heads, length, head_dim, element_bytes):
return 2 * batch * layers * kv_heads * length * head_dim * element_bytes
short = buffers(2, 16, 2048, 64, 2)
long = buffers(2, 16, 4096, 64, 2)
assert long["scores"] == 4 * short["scores"]
assert long["qkv"] == 2 * short["qkv"]
assert short["scores"] == 256 * MIB
assert short["qkv"] == 24 * MIB
for length in (2048, 4096):
sizes = buffers(2, 16, length, 64, 2)
cache = kv_cache(2, 24, 16, length, 64, 2)
print(length, sizes["scores"] / MIB, sizes["qkv"] / MIB, cache / MIB)
assert kv_cache(2, 24, 4, 2048, 64, 2) * 4 == kv_cache(2, 24, 16, 2048, 64, 2)Ожидаемые строки: 2048 256.0 24.0 384.0 и 4096 1024.0 48.0 768.0. После длины идут размеры полной матрицы оценок одного слоя, суммы Q/K/V одного слоя и KV-cache всех 24 слоёв, в MiB. Складывать эти три колонки как универсальный пик памяти нельзя: это разные объекты, их одновременная жизнь зависит от режима исполнения.
KV-cache хранит K и V уже обработанных токенов. В его формуле нет второго N, зато есть количество слоёв. Если архитектура использует меньше KV-голов, чем query-голов, в расчёте cache нужно брать именно KV-головы. Последний assert проверяет влияние этой замены в нашей идеализированной формуле, не изменение качества модели.
Посмотрите на настоящее выделение небольшого буфера
Следующий отдельный файл создаёт только компактные массивы CPU. У типа array('f') размер элемента берётся из itemsize, а не угадывается по объекту Python float. tracemalloc запускается заново для каждого замера и фиксирует пик прослеживаемых выделений Python; он не измеряет память GPU или всю память процесса.
from array import array
import tracemalloc
def measure(elements):
tracemalloc.start()
values = array("f", [0.0]) * elements
payload = len(values) * values.itemsize
current, peak = tracemalloc.get_traced_memory()
tracemalloc.stop()
return payload, current, peak
small = measure(256 * 256)
large = measure(512 * 512)
tile = measure(64 * 64)
assert large[0] == 4 * small[0]
assert tile[0] * 16 == small[0]
for name, result in (("N=256", small), ("N=512", large), ("tile=64", tile)):
print(name, "payload/current/peak bytes:", result)Первая колонка каждого результата — полезные байты массива. На платформе с четырёхбайтовым элементом это 262144, 1048576 и 16384 байта соответственно. Текущий расход и пик могут быть немного больше из-за служебных объектов; их точное значение зависит от среды. Проверки намеренно относятся к размерам данных, а не к случайному числу байтов заголовка объекта.
Размер одного tile не равен всей памяти блочного attention. Нужны Q/K/V, выход, статистики нормировки и другие буферы; кроме того, блоков много. Экономия возможна, когда промежуточные блоки используются и освобождаются по очереди. Если сохранить все tiles в список, полный расход вернётся. И один маленький массив не реализует attention: данный опыт показывает только геометрию хранения.
Локализуйте причину нехватки памяти
Для собственной модели сравните два запуска с одинаковыми весами, dtype и batch, меняя только длину. Отдельно измерьте обработку исходного контекста и генерацию новых токенов. Запишите размер входа, число выходных токенов и пик памяти для каждой фазы. Если растёт cache, оптимизация временной матрицы оценок может не устранить ограничение на число одновременно обслуживаемых запросов.
При обучении нужен отдельный замер с backward и состоянием оптимизатора. Показатель инференса без градиентов нельзя подписывать как «память обучения». Также различайте память занятых тензоров и резерв аллокатора. Расчёт формы массива — ориентир для постановки вопроса, а показания профилировщика — наблюдение конкретного запуска.
Causal mask ограничивает доступ к будущим позициям, но сама по себе не гарантирует половинный размер выделения: реализация может всё равно использовать квадратный массив. Аналогично, название функции не доказывает выбор эффективного ядра. Для отчёта зафиксируйте фактически выполненный backend и сравните одинаковую математическую задачу.
Проверьте себя изменением условий
Добавьте в калькулятор оценку mask из одного байта на элемент и выходного массива. Затем сравните три изменения по отдельности: удвоение batch, удвоение контекста и удвоение числа KV-голов. Для каждого назовите затронутые буферы. Не делайте вывод о KV-cache по числу query-голов, если они различаются.
Итог практикума — короткий отчёт с формами массивов, предположениями о dtype, таблицей расчётов и одним измерением. Если расчёт и наблюдение расходятся, сначала перечислите недостающие буферы и время их жизни. После этого можно проверять изменение реализации attention, лимитов контекста или размера batch по одному параметру за запуск.