Назад к блогу

JoLT: совместное распределение ранга и бит в KV-кэше

JoLT: совместное распределение ранга и бит в KV-кэше

Исследователи предлагают JoLT — метод сжатия KV-кэша, который распределяет бюджет хранения между рангом низкорангового разложения и точностью квантования совместно, а не по отдельности. Такой подход позволяет ужать кэш в 2–3 раза почти без потери качества и без переобучения модели. В статье разбирается механика метода, включая работу с pre-RoPE ключами и выбор ядра внимания.

KV-кэш — KV-кэш — главное узкое место по памяти при инференсе длинного контекста. Существующие методы сжатия применяют низкоранговую факторизацию или квантизацию по отдельности, не распределяя совместно ранг и точность в рамках общего бюджета хранения. JoLT рассматривает сгруппированные prefill-кеши как тензоры четвёртого порядка и применяет частичное разложение Такера по токенному и признаковому режимам, оставляя режимы голов и слоёв нетронутыми. Ниже — механика: как устроена базовая арифметика памяти, как работает трекер, как выбиралось ядро внимания и что показывают замеры.

Сколько памяти занимает KV-кэш

Объём KV-кэша считается по формуле:

KV-Cache = 2 × num_layers × num_kv_heads × d_head × seq_len × dtype_bytes

Множитель 2 — это пары K и V, num_layers — число слоёв, num_kv_heads — число голов ключей/значений, d_head — размерность одной головы, seq_len — длина контекста, dtype_bytes — байт на элемент.

Для Llama 3.1 70B заданы num_layers = 80, num_kv_heads = 8, d_head = 128, dtype = BF16 (2 байта). Подстановка даёт 2 × 80 × 8 × 128 × 2 = 327,680 байт ≈ 320 КБ на один токен. Масштабирование на длину контекста даёт 1.25 ГБ при 4,096 токенах и 10 ГБ при 32,768 токенах. При FP8 оценки вдвое меньше: ~160 КБ на токен и ~20 ГБ на 128K.

Без GQA (GQA) в Llama 3.1 70B было бы 64 KV-головы, и на один токен выходит 2 × 80 × 64 × 128 × 2 = 2,621,440 байт ≈ 2.5 МБ, а на 128K контекст — 320 ГБ, что не влезает ни на один GPU. GQA использует 8 KV-голов вместо 64, каждая KV-голова обслуживает 8 query-голов, что сокращает KV-кэш ровно в 8 раз. Падает именно множитель num_kv_heads с 64 до 8, остальные (2, num_layers=80, d_head=128, dtype_bytes=2) не меняются. На 128K контексте с GQA кэш составляет 40 ГБ вместо 320 ГБ.

Что значит «joint rank-bit allocation»

«Joint rank-bit allocation» — это совместное распределение ранга r и числа бит на элемент под общим бюджетом хранения. Две оси сжатия — ранг r (сколько компонент низкорангового разложения оставляется) и число бит на элемент (точность квантования остатка) — распределяются не по отдельности, а одним общим механизмом: повёрнутый низкобитный квантователь захватывает остаток от усечения, а один лагранжев двойственный множитель распределяет ранги Такера и битовые ширины остатка по группам под глобальным ограничением на байты.

Результат называют «near-lossless», потому что при сжатии 2–3x деградация качества мала: JoLT достигает 2–3x сжатия с менее чем 0.2% деградации перплексии без переобучения, а на RULER при контексте 64K точность поиска остаётся near-lossless до 3x и падает лишь на 0.90 и 2.40pp при 4x и 5x соответственно.

Механика трекера: pre-RoPE ключи

Инкрементальный SVD-шаг живёт в src/kvdlra/tracker/ вместе с альтернативами Oja и frequent-directions. Ключи факторизуются до RoPE: хранимый ключ имеет вид K_t = R_t U c_t, где R_t — по-токенная ротация, U — базис, c_t — координаты.

Ротация не коммутирует с проекцией, поэтому q·K_t != (Uᵀq)·c_t, и приходится восстанавливать весь n-мерный middle: кэш строит fp32 U·C, заново ротирует в истинных позициях и кастует в bf16, а при чтении склеивает в плотный массив n x T, который attention читает каждый шаг. В этот массив входят четыре яруса: sinks, exact tier, gist и recent ring. Это даёт Theta(T) токенов в n измерениях, поэтому отношение хранения не доходит ни до памяти, ни до пропускной способности.

Рассматривались альтернативы. Вариант (i) — трекать уже повёрнутые ключи K_t ≈ U c_t, тогда q_h·K_t = (U[kv(h)]ᵀ q_h)·c_t, ничего n-мерного не формируется и RoPE полностью уходит из decode. Вариант (ii) — факторизация по RoPE-парам: q_h · R_t U c_t = Σ_j [cos θ_j t, sin θ_j t] · (Q_{h,j}ᵀ U_{h,j} c_t), где A_{h,j} = Q_{h,j}ᵀ U_{h,j} строится один раз на слой за шаг, а на токен идёт 2×r matvec на (head, pair). Решение — вариант (iii): он сохраняет ровно ту же пред-RoPE конструкцию, что и в проде, двигает те же байты, что (i) и (ii), и его дополнительная стоимость — только FLOPs.

Сравнение альтернатив трекера

Сравнение с основным методом (isvd) ведётся по критерию TOST с допуском ±0.02 бит/токен и по точности на задачах NIAH (niah_single, niah_multikey, niah_multivalue) с порогом не-ухудшения delta ≤ 0.03 и Holm p < 0.05. В таблице gate1.md для llama сравнение isvd vs fd по TOST не проходит (d=-0.0741 бит, p=1), а для qwen — isvd vs frozen не проходит (d=-0.0137 бит, p=0.09) и isvd vs fd не проходит (d=-0.0939 бит, p=1). При этом qwen/niah_multikey показывает разделение isvd vs fd (Holm p=0.027).

Выбор ядра внимания и cost model

В ADR рассматриваются три варианта factored-attention kernel: (i) post-RoPE tracking at higher rank, (ii) RoPE-pair-factored pre-RoPE kernel и (iii) Triton tile-wise reconstruct inside attention.

Для (i) стоимость на низкоранговый токен, слой и последовательность составляет 4·H_q·r FLOPs (8,192 при r=64) и трафик 2·r·b_C байт; в таблице — 8,192 FLOPs/token, 512 B чтения, 0 записи. Для (ii) общая стоимость 532,480 FLOPs/token при r=64, что ровно в 4 раза больше ключевой части (iii), и трафик 516 B чтения, 0 записи. Для (iii) стоимость 4·n·r + 3n + 4·H_q·D = 281,600 FLOPs при r=64 и тот же трафик 2·r·b_C + 4 байта, что в таблице даёт 516 B чтения и 0 записи.

Варианты (i), (ii) и (iii) побайтово идентичны по HBM-трафику и резидентной памяти — все читают U один раз на слой, стримят C, читают плотные tiers — и хранят одно и то же состояние, поэтому обе метрики сводятся к одной паре столбцов на ранг; различаются они только FLOPs.

Почему выбран вариант (iii)

Вариант (iii) сохраняет уже отгруженный pre-RoPE дизайн: он сохраняет ровно ту же пред-RoPE конструкцию, что и в проде, двигает те же байты, что (i) и (ii), и его дополнительная стоимость — только FLOPs. В отличие от (i), который отказывается от pre-RoPE operating point, (iii) оставляет базис без изменений. Байты, которые он двигает, — это те же 2·r·b_C + 4 на low-rank токен, где +4 — это int32 истинная позиция, уже хранимая в mid_pos и уже учтённая в accounting.bug_footprint.

В cost model заложены константы: bf16 = 2 B, fp32 = 4 B, U и C fp32 at rest; A100-40GB HBM 1.555 TB/s и 312 TFLOP/s bf16 dense. Для (iii) r64 таблица даёт 281,600 FLOPs/token, 516 B/token чтения HBM и 0 записи. SRAM на программу для (iii) блокируется как один KV head × M токенов, с U_k,U_v head slices 2·D·r·2 B резидентно на протяжении tile loop.

Почему cos/sin строят в ядре

Стримить таблицу cos/sin в tile нельзя: при 64K это 67 МБ на шаг — слишком много трафика. Вместо этого передают собственный inv_freq модели (D/2 fp32, уже несущий rope scaling Llama-3.1) и attention_scaling, а cos/sin строят в fp32 прямо в ядре из int32-позиций тайла. fp32 обязателен, а не предпочтителен, так как углы достигают ~10⁵ рад на 128K. Стоимость — ~3n FLOPs/токен, ~1% работы тайла. Поскольку (i), (ii) и (iii) байт-идентичны и читают U один раз на слой, стримят C и читают плотные tiers, замена таблицы на построение в ядре убирает лишний трафик, оставляя только FLOPs.

Конфигурации pods и гиперпараметры

Имена вида r64-h256 и r128-h1024-s32 задают три числа: r — ранг низкоранговой части, h — размер точного яруса (exact tier), s — размер недавнего кольца (recent ring). В ADR для r64 сказано: конфигурация r64 держит 4 verbatim sink, exact tier 256, recent ring на high water 47, absorb block 16; T_mid = T − 307 low-rank columns. То есть h=256 — это exact tier, а s=32 в имени r128-h1024-s32 — recent ring.

exact tier — это h токенов, recent ring — s токенов, verbatim sinks — первые токены последовательности, а absorb block задаёт, какими порциями обрабатываются низкоранговые колонки.

Матрицы U (n×r) и C (r×T_mid) хранятся в fp32 at rest и общие для всех KV-голов слоя. Итоговый объём хранения выражен как stored state: для bugSseed-r64-h256 он равен 0.149x у Qwen2.5-7B и 0.085x у Mistral-7B-v0.3 и Llama-3.1-8B, а для bugSseed-r128-h1024-s32 — 0.284x. В таблице 3 stored state определён как fp32-at-rest stored bits (sbits=), а в таблице 1 — как float-equivalent ratio (ratio=).

Как сравнивали и что получилось

Стенд и методология

Эксперименты зафиксированы на A100-40GB с HBM 1.555 TB/s и 312 TFLOP/s bf16 dense; bf16 считается как 2 байта, fp32 — как 4 байта, а U и C хранятся в fp32 at rest. Память учитывается как resident KV (stored gist + tiers) плюс transient workspace, при этом fp32 U и C считаются по 32 бита, verbatim sinks/ring/tier — по 16, aux words — по 32. Для сравнения участников заданы строки full KV, reconstruct r64, (i) r64/r128, (ii) r64 и (iii) r64/r128; (i), (ii) и (iii) побайтово идентичны и различаются только FLOPs. Ограничения SRAM: на A100 164 KB/SM, на H100 228 KB/SM; r=64 с M=32 — единственная конфигурация, достигающая 2 blocks/SM на A100, а r=128 с M=128 не помещается на A100.

В кросс-модельный набор входят Qwen2.5-7B, Mistral-7B-v0.3 и Llama-3.1-8B. Table 1 — cross-model 16K retrieval, config bugSseed-r64-h256; Table 2 — cross-model 32K retrieval, config bugSseed-r64-h256; Table 3 — 32K variable-tracking на Llama-3.1-8B, config bugSseed-r128-h1024-s32 vs baselines. Число trials видно через знаменатели ячеек: в Table 1 и Table 2 это n = 12, в Table 3 у конфига bugSseed-r128-h1024-s32 — n = 16, у quant-4bit-kivi и bugSseed-r256-h1024 — n = 12.

Кросс-модельные таблицы (16K и 32K)

Table 1 и Table 2 показывают точность (acc) с интервалом Уилсона 95% и числом попаданий (hits/n) для трёх моделей на задачах single, multi-key, multi-value и var-track, а также stored state.

На 16K (Table 1) все три модели дают 1.00 на single, multi-key и multi-value; на var-track Qwen2.5-7B даёт 1.00, Mistral-7B-v0.3 — 0.50 (6/12), Llama-3.1-8B — 0.58 (7/12). Stored state на 16K: Qwen2.5-7B 0.149x, Mistral-7B-v0.3 и Llama-3.1-8B по 0.085x.

На 32K (Table 2) single и multi-key у всех 1.00; multi-value у Mistral-7B-v0.3 падает до 0.83 (10/12), у Qwen2.5-7B и Llama-3.1-8B остаётся 1.00; var-track: Qwen2.5-7B 1.00, Mistral-7B-v0.3 0.42 (5/12), Llama-3.1-8B 0.92 (11/12). Stored state на 32K: Qwen2.5-7B 0.139x, Mistral-7B-v0.3 и Llama-3.1-8B по 0.075x.

Сравнение с KIVI при matched bytes (Table 7)

Table 7 (KIVI 2/4-bit при matched bytes) даёт по моделям и контекстам:

Модельctxarmstoredsinglemulti-keymulti-valuevar-track
Llama-3.1-8B16384r640.151x1.001.001.000.58
Llama-3.1-8B163842-bit0.163x1.000.670.420.67
Llama-3.1-8B163844-bit0.287x1.001.001.001.00
Llama-3.1-8B32768r640.139x1.001.001.000.92
Llama-3.1-8B327682-bit0.160x1.000.830.920.92
Llama-3.1-8B327684-bit0.284x1.001.001.001.00
Mistral-7B-v0.316384r640.150x1.001.001.000.50
Mistral-7B-v0.3163842-bit0.163x0.920.580.500.33
Mistral-7B-v0.3163844-bit0.287x1.001.001.000.25
Mistral-7B-v0.332768r640.139x1.001.000.830.42
Mistral-7B-v0.3327682-bit0.160x0.830.250.080.00
Mistral-7B-v0.3327684-bit0.284x1.001.001.000.25
Qwen2.5-7B16384r640.275x1.001.001.001.00
Qwen2.5-7B163842-bit0.163x1.000.830.330.92
Qwen2.5-7B163844-bit0.287x1.001.000.921.00
Qwen2.5-7B32768r640.265x1.001.001.001.00
Qwen2.5-7B327682-bit0.160x0.920.580.170.25
Qwen2.5-7B327684-bit0.284x1.001.001.001.00

32K variable-tracking на Llama-3.1-8B (Table 3)

Конфигурация bugSseed-r128-h1024-s32 даёт var-track 0.94 [0.72,0.99] (15/16) при stored state 0.284x. Для сравнения: think-c0.5 — 0.31 (5/16) при 0.750x, p=2.0e-03, discordant 10/0; palu-r0.5 — 0.56 (9/16) при 0.502x, p=3.1e-02, discordant 6/0; quant-4bit-kivi — 1.00 (12/12) при 0.284x, p=1.0e+00, discordant 0/1; bugSseed-r256-h1024 — 0.00 (0/12) при 0.534x, p=9.8e-04, discordant 11/0.

NVIDIA RULER 16K на Llama-3.1-8B (Table 8)

Mean по девяти задачам: full 0.99, think-c0.5 0.98, palu-r0.5 0.92, quant-4bit-kivi 0.99, quant-2bit-kivi 0.87, bugSseed-r64-h256 0.79, bugSseed-r64-h256-q4 0.00, ea-k0.1 0.20.

Eviction до 0.1x и 0.25x

При eviction до 0.1x от сохранённого состояния (arm ea-k0.1) первыми падают метрики var-track и multi-value/multi-key у Qwen2.5-7B: на 16384 все три равны 0.00 (0/12), а single — 0.17 (2/12). У Mistral-7B-v0.3 на 16384 single держится на 0.50 (6/12), тогда как multi-key 0.08 (1/12), а multi-value и var-track — 0.00 (0/12). У Llama-3.1-8B на 16384 single и multi-value остаются 1.00 (12/12), multi-key 0.92 (11/12), но var-track падает до 0.08 (1/12); на 32768 var-track восстанавливается до 0.58 (7/12).

При менее агрессивном eviction до 0.25x (arm ea-k0.25) метрики заметно выше: single 1.00, multi-key 0.88, multi-value 1.00, var-track 0.50. Для сравнения, полный кэш представлен baseline shadow-r64 со stored state 0.815x, где single и multi-key равны 1.00, а var-track — 0.00.

Gate 1: bf16 non-inferiority

В Gate 1 для семейств llama и qwen при ctx 16384 измеряются три задачи извлечения: niah_single, niah_multikey и niah_multivalue. В каждой ячейке стоит delta [a_favored/b_favored of n_paired] Holm p, где a — это isvd_r64_h256_seed, а b — isvd_r64_h256_seed_bf16, поэтому delta > 0 означает, что пара проиграна bf16-вариантом. Задача считается не-инфериорной, если не выполняется одновременно delta > 0.03 и Holm p < 0.05.

Для llama при ctx 16384: niah_single +0.042 [2/1 of 24] Holm p 1, niah_multikey +0.125 [3/0 of 24] Holm p 1, niah_multivalue +0.042 [1/0 of 24] Holm p 1; TOST проходит со значением -0.0007 bits, sbits bf16/isvd = 0.5667, вердикт PASS. Для qwen при ctx 16384: niah_single +0.000 [0/0 of 24] Holm p 1, niah_multikey +0.083 [2/0 of 24] Holm p 1, niah_multivalue +0.042 [1/0 of 24] Holm p 1; TOST проходит со значением +0.0102 bits, sbits = 0.5382, вердикт PASS. Holm-коррекция применяется при alpha=0.05 внутри собственного семейства извлечения при реализованном m=6.

Итог по bf16: llama PASS, qwen PASS. Общий вердикт — UNDECIDED: ничего не выбрано, поскольку для llama TOST isvd против fd при ±0.02 не проходит (d=-0.0741 бит, p=1), для qwen TOST isvd против frozen не проходит (d=-0.0137 бит, p=0.09), qwen/niah_multikey разделяется с fd (Holm p=0.027), и qwen TOST isvd против fd не проходит (d=-0.0939 бит, p=1).

Что цифры не показывают

Числа в таблицах — это точечные оценки точности (acc) с интервалами Уилсона, а не измерения задержки или пропускной способности: ячейка — это acc [Wilson 95% lo,hi] (hits/n), и никаких полей latency/throughput в таблицах нет. Опущенные строки v1 несут только архивную acc без записей по отдельным прогонам, поэтому n неизвестен и интервал Уилсона не печатается. Опущены были две конкурентные строки baselines — ShadowKV после исправления (shadow-r64) и eviction при 0.25x (ea-k0.25). В архивных строках модель записана как unknown, а атрибуция к Llama-3.1-8B сделана в CODE_AUDIT.

Воспроизводимость

Устройство пайплайна

«Pod» — это один GPU-эксперимент: файл configs/pods/<pod>.yaml, в котором перечислены его arms, tasks, model и image, плюс prereg/<pod>.md, где сказано, что он будет запускать и какой исход что будет означать. Порядок обязателен: сначала коммит prereg и конфига, затем launch, который отказывается работать, если prereg-коммит не является строгим предком HEAD, затем harvest лога в records и check.

make env создаёт окружение: uv venv + CPU torch + uv pip install -e ".[dev]". Числа таблиц регенерируются из закоммиченных per-trial records командой make tables, которая пишет docs/paper/tables/ из results/paper-v1/*/trials.jsonl и сверяет результат с docs/plan/paper-v1-tables.md, печатая tables: diff-clean или падая. make figures пересобирает три фигуры из тех же records, а make check заново проверяет каждый закоммиченный каталог results.

Требования и структура репозитория

Окружение требует Python ≥ 3.11 и uv; CI запускается на 3.12. Структура репозитория включает src/kvdlra/ с подкаталогами tracker/ (инкрементальный SVD-шаг и альтернативы Oja / frequent-directions), cache/ (потоковые кэши: low-rank cache, ShadowKV), baselines/, quant/, eval/ и accounting.py. Каталог configs/ содержит arms/, tasks/, pods/ — по одному YAML на arm, task и pod, так что каждый эксперимент является конфигом. Каталог docs/paper/tables/ заполняется командой make tables.

Как повторить конкретную таблицу

Сначала клонируют репозиторий и ставят окружение: git clone https://github.com/hkrishna42/kvdlra.git && cd kvdlra, затем make env, затем make tables. Для Table 1 источником служат results/paper-v1/w18-g1-{qwen,mistral,llama}/trials.jsonl (ячейки) и results/paper-v1/w18-{qwen,mistral,llama}/cells.jsonl (stored state). Проверка совпадения чисел с опубликованными — это diff против docs/plan/paper-v1-tables.md; дополнительно make check перепроверяет каждый закоммиченный каталог результатов. Запуск пода идёт через scripts/pod.py launch --pod mypod --offer <vast-offer-id>, затем scripts/pod.py harvest --pod mypod и scripts/pod.py check results/mypod. check требует, чтобы каждая ячейка конфига (arm × generator × sub-task × context) содержала ровно n_trials × len(seeds) записей.

Что из этого следует на практике

  • Совместное распределение ранга и бит — не то же самое, что их независимый подбор. Один лагранжев двойственный множитель распределяет ранги Такера и битовые ширины остатка по группам под глобальным ограничением на байты, тогда как существующие методы применяют факторизацию или квантизацию по отдельности.
  • Pre-RoPE факторизация требует восстановления n-мерного middle. Ротация не коммутирует с проекцией, поэтому q·K_t != (Uᵀq)·c_t, и без явного восстановления attention обрабатывает Theta(T) токенов в n измерениях, из-за чего отношение хранения не доходит ни до памяти, ни до пропускной способности.
  • Выбор ядра определяется не трафиком, а FLOPs. Варианты (i), (ii) и (iii) побайтово идентичны по HBM-трафику и резидентной памяти и различаются только FLOPs; (iii) выбран потому, что сохраняет уже отгруженный pre-RoPE дизайн, а его дополнительная стоимость — только FLOPs.
  • Таблицу cos/sin нельзя стримить. При 64K это 67 МБ на шаг; вместо этого передают inv_freq модели и attention_scaling, а cos/sin строят в fp32 в ядре из int32-позиций тайла, поскольку углы достигают ~10⁵ рад на 128K.
  • Деградация при агрессивном eviction неравномерна. При eviction до 0.250x первыми падают var-track и multi-value/multi-key; у Qwen2.5-7B на 16384 все три равны 0.00, тогда как у Llama-3.1-8B single и multi-value остаются 1.00, а var-track падает до 0.08.
  • Числа в таблицах — только точность. Это acc с интервалами Уилсона, а не замеры задержки или пропускной способности; опущенные строки baselines несут только архивную acc без per-trial записей, поэтому n неизвестен и интервал не печатается.

Где смотреть в коде

Источники

Похожее