Перейти к содержимому
13/30Глава 13 из 30

Как удешевить inference: KV Cache, batching и quantization

Один и тот же model отвечает за 8,8 и 78,9 секунды с byte-identical выводом. Затем INT4: три измерения вместо заявлений.

На этой странице

Один и тот же model, на той же машине, отвечает на тот же вопрос теми же 48 token. Два вывода идентичны token за token — это проверено, а не предположено.

TEXT
with a key-value cache:     8.85 s   ( 6.01 tokens/second)
without a key-value cache: 78.95 s   ( 0.60 tokens/second)

Изменился один аргумент: use_cache=False. Ничего в model, prompt, sampling или арифметике не стало другим, и второй запуск не стал точнее в обмен на свои усилия. Он просто в девять раз медленнее без причины.

Такова форма этой главы. Всё в ней — cache, batch, quantized weights — это попытка перестать платить за работу, которая не меняет ответ, или выяснить, сколько стоит более дешёвый ответ. Глава 10 установила прайс-лист для обучения. Это прайс-лист для стороны, за которую вы платите всегда: развёрнутый model тратит примерно 2N2N FLOPs на каждый token, который выдаёт, в каждом запросе, до конца своей жизни.

Чтобы сгенерировать token, decoder-only transformer берёт всю последовательность на данный момент, прогоняет её через каждый слой и считывает распределение вероятностей с последней позиции. Затем добавляет выбранный token и делает это снова. Это описание корректно, и именно это делает медленный запуск.

Оно также чрезвычайно расточительно, и причина — causal mask из главы 9. Векторы key и value позиции 7 вычисляются из входа позиции 7 и позиций перед ней. Когда приходит позиция 8, позиция 7 не может её видеть — именно это и означает causal, — поэтому key и value позиции 7 являются ровно теми же числами, что и раньше. Медленный запуск всё равно пересчитывает их на каждом шаге.

Значит, сохраните их. Это хранилище и есть key-value cache, самая важная оптимизация в serving языковых model:

generate.pyPYTHON
out = model(prompt_ids, use_cache=True)          # prefill: the whole prompt
past = out.past_key_values                        
nxt = out.logits[:, -1].argmax(-1, keepdim=True)

for _ in range(n - 1):
    out = model(nxt, past_key_values=past, use_cache=True)   
    past = out.past_key_values                                
    nxt = out.logits[:, -1].argmax(-1, keepdim=True)

Посмотрите, что подаётся в model внутри цикла: nxt, один token. Не последовательность. Query нового token attends against каждый cached key, а cached keys всё равно никогда не должны были измениться. Это не approximation — проверка идентичности вывода выше и есть суть. Cache не обменивает качество на скорость; он удаляет лишнюю арифметику.

Чтобы ясно увидеть scaling, уберите transformer и замерьте один attention head с d=64d = 64, один шаг generation, посчитанный обоими способами:

token в contextпересчитать всёс cacheотношениеscore matrix
1280,59 ms0,062 ms10x65 536 B vs 512 B
2561,20 ms0,163 ms7x262 144 B vs 1 024 B
5127,03 ms0,078 ms90x1 048 576 B vs 2 048 B
102417,31 ms0,114 ms152x4 194 304 B vs 4 096 B
204859,83 ms0,214 ms279x16 777 216 B vs 8 192 B
4096236,18 ms0,284 ms832x67 108 864 B vs 16 384 B

Правый столбец — причина. Пересчёт строит полную n×nn \times n attention matrix на каждом шаге — ту самую O(n2)O(n^2) из блока asymptotic notation в главе 9, оплачиваемую один раз на token. С cache вы вместо этого строите строку 1×n1 \times n: при 4 096 token — 67 MB scores против 16 KB.

Если считать multiply-accumulates вместо миллисекунд, машина исчезает из аргумента. Чтобы сгенерировать TT token с cold start:

сгенерировано tokenс cacheпересчётотношение
1282,6 M192,0 M73x
51223,1 M7,36 G318x
2048293,7 M392,6 G1 336x

На каждом шаге версия с cache линейна по context, а без cache — квадратична; если суммировать по generation, получается O(T2)O(T^2) против O(T3)O(T^3), и отношение растёт без предела. Девятикратная разница в начале была измерена всего на 48 token — меньше первой строки этой таблицы.

Cache также меняет то, что должно находиться в памяти. На laptop GPU с 8 GB при generation 256 token в fp16, если взять пик allocator и вычесть resident weights:

peak working memory
с cache21,8 MB
пересчёт181,7 MB

В 8,3 раза больше памяти, потраченной на то, чтобы произвести те же token медленнее. Это обещание из главы 5, пришедшее с неожиданной стороны: там reverse-mode autodiff должен был держать каждый intermediate живым для backward pass, и activations доминировали в памяти обучения. В inference нет backward pass и нечего удерживать ради него — поэтому память вместо этого доминируется cache, и это осознанный выбор, а не неизбежная стоимость.

Снова посмотрите на быстрый запуск: его первый token вёл себя не так, как остальные сорок семь.

TEXT
prefill, 40 prompt tokens : 1.0224 s   ->  25.6 ms per token
decode,  47 steps         : 0.1665 s mean per step

Prompt стоил 25,6 ms на token, а каждый сгенерированный token — 166 ms. Тот же model, то же hardware, те же weights, шестикратная разница на token — и в направлении, которого большинство не ожидает. Prompt — дешёвая часть. Generation распадается на две фазы с действительно разной физикой:

Один forward pass по всему prompt. Каждый token обрабатывается параллельно, поэтому каждая weight matrix загружается из памяти один раз и умножается на matrix из сотен token vectors — matrix-matrix product, с большим количеством арифметики на каждый перемещённый byte, то есть именно то, для чего создан GPU. Prefill является compute-bound, и его стоимость примерно линейна по длине prompt.

Один forward pass на token, batch из одного и sequence из одного. Каждая weight matrix всё ещё полностью загружается из памяти и умножается на один vector — matrix-vector product, почти без арифметики на перемещённый byte. Decode является memory-bandwidth-bound, и его стоимость на token почти не зависит от длины context.

Обе половины измеримы. Prefill, один pass по PP token:

prompt tokenssecondsms per token
160,351521,97
320,525416,42
641,049116,39
1281,655212,93
2563,096512,10

Decode, один token against cache из CC:

cached tokensms for one token
16110,05
6497,57
256108,53
1024103,86

Прочитайте вторую таблицу дважды. Переход от 16 token context к 1 024 — в шестьдесят четыре раза больше history для attention — не изменил стоимость шага измеримо. Attention against cache — реальная работа, но её затмевает фиксированная стоимость протаскивания полумиллиарда weights через memory bus, чтобы получить один vector. Эта фиксированная стоимость — причина всего в следующем разделе.

Эти две фазы — источник двух чисел, которые сообщает любая serving system. Time to first token — по сути prefill, и он растёт вместе с prompt, поэтому длинный разговор медленно стартует. Tokens per second — это 1/decode step1/\text{decode step}, и он примерно постоянен, поэтому ответ затем течёт ровно. Chat, который медленно начинается, а потом плавно streamится, — не трюк rendering. Это эти две таблицы.

Cache обменивает арифметику на память, и память, которую он хочет, не мала. Для каждого token в context каждый layer хранит один key vector и один value vector на key-value head:

bytes per token=2×L×Hkv×dhead×bytes per element\text{bytes per token} = 2 \times L \times H_{kv} \times d_{\text{head}} \times \text{bytes per element}

2 — это keys и values; всё остальное — architecture. Для model, измеряемого во всей этой главе, — 24 layers, 14 query heads, 2 key-value heads, head dimension 64 — в fp16 это 2×24×2×64×2=12,2882 \times 24 \times 2 \times 64 \times 2 = 12{,}288 bytes на token.

Формулы в этой области имеют привычку ошибаться на factor of two, поэтому проверьте её по allocator, а не верьте на слово:

TEXT
KV cache tensors per layer: (1, 2, 295, 64) float16
measured: 3,624,960 bytes for 295 tokens = 12,288 bytes/token
formula : 2 * 24 * 2 * 64 * 2                = 12,288 bytes/token

Точно, и остаётся точно для каждой проверенной shape:

batchcontextmeasured cachepredictedpeak working memory
15126,0 MB6,0 MB15,4 MB
116 384192,0 MB192,0 MB207,3 MB
165 536768,0 MB768,0 MB793,7 MB
84 096384,0 MB384,0 MB401,5 MB
322 048768,0 MB768,0 MB794,2 MB
641 024768,0 MB768,0 MB797,0 MB
128512768,0 MB768,0 MB816,4 MB

На последние три строки стоит посмотреть ещё раз. Тридцать два пользователя по 2 048 token у каждого, шестьдесят четыре по 1 024, сто двадцать восемь по 512 — cache во всех случаях равен 768 MB, потому что во всех трёх находятся 65 536 token. Cache зависит только от общего числа resident token, а не от того, как они распределены между пользователями. Этот факт — фундамент раздела про batching.

Глава 9 представила multi-query и grouped-query attention и отложила причину до этой главы. Причина — эта формула, а точнее HkvH_{kv} в ней.

Стандартный multi-head attention даёт каждому query head собственные key и value heads. У этого model 14 query heads; с полным multi-head attention его cache был бы 2×24×14×64×2=86,0162 \times 24 \times 14 \times 64 \times 2 = 86{,}016 bytes на token — 84 KB вместо 12 KB, ровно в семь раз больше, по отношению query heads к key-value heads.

Multi-query attention1 доводит это до предела: все query heads делят один key-value head. Grouped-query attention2 — компромисс, который победил: несколько key-value heads, каждый разделяется группой query heads, — потому что потеря качества у MQA была реальной, а у GQA нет. Ни один из них не покупает арифметику. Они существуют, чтобы делить эту формулу на целое число, и распространились по индустрии в тот момент, когда длинные contexts сделали cache binding constraint.

А он делает это быстро. Для model класса 7B с 32 layers и 8 key-value heads размерности 128 cache составляет 128 KB на token в fp16:

context tokensone user8 users64 users
4 0000,49 GB3,91 GB31,2 GB
32 0003,91 GB31,25 GB250,0 GB
128 00015,62 GB125,00 GB1 000,0 GB
1 000 000122,07 GB976,56 GB7 812,5 GB

Собственные weights этого model — 13,0 GB в fp16, число из таблицы в конце главы. Поэтому при context в 128 000 token cache одного пользователя больше, чем model. Эту арифметику глава 16 превращает в деньги, и именно поэтому длинный разговор не просто медленный — он занимает фиксированный кусок машины всё время, пока жив запрос.

Batching: число, которое растёт, и число, которое падаёт

Ссылка на раздел: Batching: число, которое растёт, и число, которое падаёт

Decode memory-bound: weights протаскиваются через bus, чтобы произвести один token, а арифметические блоки простаивают. Значит, добавьте больше работы в тот же шаг. Запустите несколько запросов одновременно, и weights, прочитанные один раз, обслужат их всех. Измерено на том же model: каждый запрос держит cache на 64 token и decode один token:

batchlatency per stepthroughputlatency vs B=1
10,1286 s7,78 tok/s1,00x
20,1839 s10,88 tok/s1,43x
40,1909 s20,95 tok/s1,49x
80,2781 s28,76 tok/s2,16x
160,3430 s46,64 tok/s2,67x
320,6302 s50,78 tok/s4,90x

Читайте два правых столбца вместе, потому что в них весь смысл. Переход от одного запроса к шестнадцати умножает throughput на 6,0 и умножает ожидание любого отдельного запроса на 2,67. Batch сделал сервер лучше, а каждого пользователя — хуже.

Это не bug, который надо настроить до исчезновения; это сам trade-off, и у каждой стороны есть имя. Latency — то, что ощущает человек, ожидающий ответ. Throughput — то, на что делится счёт. Нет настройки, которая улучшает оба.

Обратите внимание, где это останавливается. От 16 до 32 throughput растёт на 9 %, а latency почти удваивается: шаг перестал быть memory-bound и стал compute-bound, и после этого knee batch ничего не покупает. У каждого deployment есть такой knee; его положение нужно измерять на вашем, но его существование — нет.

Static batching тратит впустую большую часть того, что выигрывает

Ссылка на раздел: Static batching тратит впустую большую часть того, что выигрывает

Наивный способ batch — собрать BB запросов, запустить их вместе и вернуть, когда все закончат. Но они не заканчивают вместе: одни ответы — двадцать token, другие — пятьсот. Fixed batch работает, пока не закончит самый длинный участник, а каждый завершённый запрос до тех пор продолжает занимать свой slot и добавлять padding.

Возьмём 64 запроса с реалистичным перекосом длин output — median 18 token, longest 231, всего 1 874 — и симулируем обе политики при измеренной per-step cost для восьми slots:

policywall clockthroughputmean latency per requestwasted slot-steps
static batches of 8176,9 s10,6 tok/s83,2 s3 214
continuous, 8 slots109,0 s17,2 tok/s8,1 s0

Throughput улучшается в 1,6x. Mean latency улучшается более чем в десять раз, потому что при static batching запрос, закончивший за четыре шага, всё равно ждёт соседа на 231 token, прежде чем кто-либо услышит ответ.

Continuous batching3 — исправление, и оно именно настолько просто, как звучит: batch — это не группа, а набор slots, и slot, который освободился, принимает следующий запрос из очереди на следующем же шаге. Scheduler работает с granularностью одного token, а не одного request. Сейчас это делает каждый production serving stack.

У него есть вторая половина — cache. Slots, которые приходят и уходят, оставляют cache memory fragmented, а reservation максимального possible context для каждого slot тратит большую часть reservation впустую. PagedAttention4 заимствует ответ у operating systems: хранить cache в fixed-size blocks с block table для каждой sequence, чтобы cache sequence мог быть физически разбросан, оставаясь логически contiguous, — что также позволяет двум sequences с shared prefix делить blocks, которые его хранят. На этом построен vLLM, и поэтому serving engine — это memory allocator с приделанным transformer.

Другая половина счёта — сами weights. Полмиллиарда parameters по четыре bytes каждый — 1,98 GB; по два bytes — 0,99 GB; по одному byte — 0,49 GB. Меньше bits на weight уменьшает model на disk, уменьшает его в memory и — поскольку decode bandwidth-bound — делает каждый шаг быстрее, потому что нужно перемещать меньше bytes.

Самая простая схема — symmetric absolute-maximum quantization, и она помещается в три строки:

quantize.pyPYTHON
qmax  = 2 ** (bits - 1) - 1
scale = W.abs().max() / qmax                        
Wq    = torch.round(W / scale).clamp(-qmax - 1, qmax)
W_hat = Wq * scale                                  # dequantized

Выберите scale так, чтобы самый большой weight отображался в самое большое integer, разделите, округлите, сохраните integers и scale. Восстановите обратным умножением. В этом нет ничего хитрого, и это работает — ровно до момента, когда перестаёт.

Измерено на настоящих weights model: все 168 projection matrices, 357,8 million parameters, relative error WW^/W\lVert W - \hat{W}\rVert / \lVert W \rVert:

schememean relative errorworst matrix
INT8, one scale for the whole matrix0,04000,1487
INT8, one scale per output row0,01000,0149
INT4, one scale for the whole matrix0,60260,9931
INT4, one scale per output row0,17900,2589
INT4, one scale per group of 1280,13230,1992
NF4, one scale per block of 640,09520,1205
INT3, one scale per group of 1280,30440,4123
INT2, one scale per group of 1280,77900,8076

Четвёртая строка — это обвал. Relative error 0,99 на worst matrix означает, что reconstruction фактически не сохраняет ничего от original — matrix заменена шумом примерно правильной magnitude. Причина видна в том же experiment на одной matrix:

TEXT
model.layers.12.mlp.down_proj.weight   (896 x 4864)
mean |w| 0.01386   std 0.01822   max |w| 0.43945   max/std 24.1
weights beyond 6 sigma: 692 of 4,358,144   (0.016 %)

Один weight из шести тысяч лежит дальше шести standard deviations, а largest — на 24. При одном scale для всей matrix этот один weight задаёт step size для всех 4,3 million остальных. При 8 bits есть 256 steps, и typical weight всё ещё попадает на meaningful one. При 4 bits их 16, крайний зарезервирован для значения, которого почти ни у чего нет, а ordinary weights — то есть почти все — округляются до двух-трёх distinct levels.

Всё после этой строки — один и тот же ремонт с разной granularity: дать scale меньшую территорию. Per output row делит error на 3,4; per group of 128 consecutive weights делит его снова. Цена — bookkeeping: 16-bit scale на group of 128 — это 4+16/128=4.1254 + 16/128 = 4.125 bits на weight вместо 4 — и это возвращает большую часть разрыва.

NF4 подходит с другой стороны.5 Levels не обязаны быть equally spaced. Weights внутри block приблизительно normally distributed, поэтому выберите шестнадцать levels как quantiles of a normal distribution: плотные near zero, где weights действительно находятся, редкие в tails, где их нет. Те же four bits, тот же block scaling, но меньший block — 4,25 bits на weight против 4,125 у group-128 — и измеренный error падает с 0,1323 до 0,0952, на 28 % ниже. Часть этого — более fine block, остальное — размещение levels там, где находится mass, а чтобы разделить эти два фактора, понадобилась бы третья строка.

Блок о floating-point в главе 2 закончился обещанием: эта глава quantize weights до 8 и 4 bits и найдёт handful of outlier features, которые откажутся сжиматься. Вот они, и они объясняют, почему «просто округлить числа» никогда не должно было сработать на activations.

Weights выше вели себя плохо. Activations — совершенно другой класс. Возьмите обычный 84-token prompt, capture residual stream на каждом layer и измерьте largest magnitude, которой достигает каждая из 896 dimensions:

layerlargest |h|median dimension's largest |h|ratiodimensions above 6x the median
16,190,33918x2
41543,481,550996x34
81571,631,4981049x36
121575,031,5461019x34
161579,601,617977x32
201577,982,361668x24
24204,4410,76019x12

Dimension 62 достигает 1 579,6, тогда как median dimension никогда не превышает 1,6. Это не случайность одного token или одного layer: та же dimension есть на layer 4 и всё ещё есть на layer 20, почти с тем же value. Это outlier features,6 и они систематичны — свойство trained model, а не input.

Histogram этих 896 per-dimension maxima на layer 16 делает форму безошибочной:

TEXT
     0 -      1 | ######################################## 254
     1 -      2 | ######################################## 283
     2 -      4 | ######################################## 226
     4 -      8 | ######################################## 93
     8 -     16 | ##################                       18
    16 -     32 | #########                                9
    32 -     64 | #######                                  7
    64 -    128 | #####                                    5
   128 -    256 |                                          0
   256 -    512 |                                          0
   512 -   1024 |                                          0
  1024 -   4096 | #                                        1

Девятьсот dimensions аккуратной кучей ниже 8, затем три октавы вообще ничего, и потом одна dimension в одиночестве на дальнем конце. Теперь quantize этот tensor в INT8 и посчитайте, что происходит:

schemerelative errordistinct integer levels used, whole tensor
one scale for the whole tensor0,108314 of 256
one scale per token (per row)0,0433158
whole tensor, 1 outlier dimension kept in fp320,044248
whole tensor, 4 outlier dimensions kept in fp320,027957
whole tensor, 16 outlier dimensions kept in fp320,0085102

Четырнадцать levels из 256. Scale был задан 1 579,6, поэтому каждый step шириной 12,44, и typical activation — median magnitude 0,26, ninety-ninth percentile 2,51 — некуда приземлиться. Per dimension это ещё резче:

TEXT
single tensor-wide scale = 12.4378
  dim 826 (max |h| = 4.77):  1 distinct level out of 256
  dim 336 (max |h| = 1.62):  1 distinct level out of 256
  dim  96 (max |h| = 0.69):  1 distinct level out of 256

after excluding the top 4 dimensions, scale = 0.5749  (22x smaller)
  dim 826: 8 levels    dim 336: 4 levels    dim  96: 3 levels

Один level. Вся dimension, каждый token, quantized в одно и то же число. Восемь bits были выделены, но использовано примерно ноль, и model, читающий эти activations, получает constant.

Это измерение — обоснование для каждой technique, которую люди действительно используют:

Не пускайте outliers внутрь. LLM.int8()6 decomposes matrix multiply: dimensions с extreme magnitudes считаются в 16 bits, всё остальное в INT8, а половины суммируются. Таблица выше — квитанция: удаление четырёх dimensions сокращает error почти в четыре раза. SmoothQuant7 вместо этого переносит сложность: разделите activations на per-channel factor и умножьте соответствующий weight column на него; product остаётся неизменным, а outlier переезжает из tensor, который не может его поглотить, в тот, который может.

Выбирайте rounding, а не просто округляйте. Ничто выше не спрашивает, для чего нужна matrix. GPTQ8 quantizes column за column и после каждой корректирует оставшиеся full-precision columns, чтобы compensate error, уже внесённый, — minimizing error output данного layer на real inputs, а не error его weights. AWQ9 замечает, что малая доля weight channels важнее остальных намного сильнее, находит их по activation statistics и масштабирует вверх перед quantizing, чтобы они попали на finer levels. Обоим нужен calibration set; ни одному не нужны gradients.

Показать детали

GGUF, и при чём здесь file format.

GGUF — не quantization method; это container, который использует llama.cpp, а путаница в сравнениях gguf vs gptq возникает из-за обращения с ними как с вещами одного рода. GGUF хранит tensors, tokenizer, architecture metadata и chat template в одном memory-mappable file и несёт внутри family block schemes — имена вроде Q4_K_M кодируют bits per weight, block size и то, оставлены ли некоторые tensors в higher precision.

Важная engineering difference: GPTQ и AWQ производят weights, optimized для GPU kernel, тогда как schemes GGUF дешёво decoded на CPU при mapped, а не loaded file. Поэтому один и тот же номинальный «4-bit 7B model» существует в обоих мирах с разными sizes и quality, и поэтому честное сравнение — никогда не format, а измерение ниже, запущенное на вашей собственной задаче.

Сколько quantization реально стоит: измерения

Ссылка на раздел: Сколько quantization реально стоит: измерения

Почти каждая статья о quantization останавливается на предыдущем разделе: объясняет method, цитирует compression ratio и утверждает, что quality «в основном сохраняется». Глава 4 была о том, как не обманывать себя, так что выясним.

Тот же model, weights quantized in place каждой scheme, затем три измерения: perplexity на 2 048 token held-out English prose — здесь это draft данного курса, поэтому repository подставляет fixed public-domain book и печатает table той же формы с другими numbers, — battery из 16 short factual questions с known answers при greedy decoding и fraction token, на которых quantized model соглашается с full-precision one при identical context.

schememean weight errorperplexityquestion batteryagrees with fp32
fp32 (reference)0,000023,0813/16100,0 %
INT8 per tensor0,040023,5813/16
INT8 per row0,010022,9613/1698,6 %
INT4 per tensor0,6026365 416 0000/16
INT4 per row0,179046,186/1658,3 %
INT4 group 1280,132331,0810/1671,5 %
NF4 block 640,095224,5511/1684,7 %
INT3 group 1280,3044213,090/165,6 %
INT2 group 1280,779026 325 4360/160,0 %

В этой таблице четыре вещи стоит сказать прямо.

Правильно сделанный INT8 бесплатен. Per-row INT8 даёт 22,96 против 23,08 у reference — разрыв в одну двухсотую, то есть noise, который следует читать как «identical». Направление noise нестабильно: на public-domain corpus из repository те же две schemes дают 22,24 против 22,18 — половина этой дистанции и в другую сторону. Он соглашается с full-precision model на 142 из 144 generated token. Четверть памяти против fp32 reference, половина против fp16, который вы реально deploy, и никакой detectible cost. INT8, сделанный небрежно, тоже почти бесплатен: один scale на matrix стоит 0,5 points perplexity и ни одного ответа battery. Восемь bits достаточно forgiving, чтобы granularity почти не имела значения, и именно поэтому люди обобщают с INT8 на INT4 и страдают.

INT4 с одним scale на tensor уничтожает model. Perplexity 365 million: не degraded, а annihilated. Затем granularity — вся игра: per-tensor 365 416 000, per-row 46,18, per-group-of-128 31,08, NF4 24,55. Те же four bits на weight, factor fifteen million между worst и best.

Perplexity — грубый инструмент, а battery ещё грубее. Между NF4 и group-128 INT4 gap perplexity равен 6,5 points, а battery отличается на один question — и confidence interval из главы 4 говорит, что один question из sixteen не различает вообще ничего. Есть демонстрация острее interval: запустите тот же battery с выключенным stock repetition penalty model — что greedy decoding actually means, — и эти две строки поменяются местами. Один question из sixteen — не small effect, а no effect. Предупреждение главы 8 тоже применимо: perplexity сопоставима только между models с общим tokenizer, поэтому число из чужого write-up нельзя сравнивать с вашим.

Столбец agreement — самый острый из трёх, и почти бесплатный: запустите full-precision model greedily, затем спросите quantized one в каждой position, что он выбрал бы при том же prefix. У него 144 independent observations вместо 16, ему не нужна ground truth, и он деградирует плавно там, где battery деградирует скачками. Это также ровно та величина, которая нужна следующему разделу.

Это обещание, которое глава 1 дала об этой главе, приходит по расписанию: mathematics говорит, что 4-bit model возможен, а engineering решает, usable ли он.

Глава 12 объявила это и оставила счёт здесь.

Идея следует напрямую из разделения prefill/decode. Проверка предложенной sequence из γ\gamma token стоит один forward pass по γ\gamma positions — matrix-matrix product, едва дороже pass по одному. Поэтому:

Маленький, дешёвый model autoregressively генерирует γ\gamma candidate token.

Большой model делает один forward pass по всем γ\gamma candidates сразу, producing то, что он сказал бы в каждой position.

Сохраните самый длинный prefix, на котором они agree, плюс token, который большой model бесплатно выдаёт при первом disagreement. Остальное отбросьте и начните снова.

Output distribution не меняется. При greedy decoding это очевидно — token принимается только если target произвёл бы его. При sampling нужна modified acceptance rule, и Leviathan et al. доказывают, что resulting distribution ровно target distribution.10 Это вторая exact optimisation в этой главе.

Следовательно, всё зависит от acceptance rate α\alpha, которую можно измерить, — это agreement column выше, поэтому он там и был посчитан. Используя каждый quantized model как draft для full-precision target, на 144 generated positions:

draft modelacceptancelongest accepted runexpected tokens per target pass, γ=4\gamma = 4
fp32 (the target itself)100,0 %485,00
INT8 per row98,6 %484,86
NF4 block 6484,7 %203,69
INT4 group 12871,5 %132,85
INT4 per row58,3 %72,24
INT3 group 1285,6 %21,06
INT2 group 1280,0 %01,00

Expected tokens accepted за verification pass при draft length γ\gamma равен

E[tokens]=1αγ+11α\mathbb{E}[\text{tokens}] = \frac{1 - \alpha^{\gamma+1}}{1 - \alpha}

а net speedup делит это на собственную стоимость draft, fraction cc от target на token:

acceptancec=0.05c=0.05, γ=4\gamma=4c=0.1c=0.1, γ=4\gamma=4c=0.2c=0.2, γ=4\gamma=4c=0.1c=0.1, γ=8\gamma=8
30 %1,19x1,02x0,79x0,79x
50 %1,61x1,38x1,08x1,11x
70 %2,31x1,98x1,54x1,78x
90 %3,41x2,93x2,28x3,40x

Bold entry — то, что нужно запомнить: speculative decoding может сделать generation медленнее. При acceptance 30 % и draft, который стоит fifth of the target, вы платите за five forward passes и сохраняете 1,4 token. Последний столбец — другая ловушка: longer draft помогает только при high acceptance, потому что tail of a γ\gamma-token guess почти никогда не достигается. При 90 % acceptance γ=8\gamma = 8 стоит 3,40x, а при 30 % — 0,79x: одна и та же configuration, win или loss в зависимости от числа, измеренного на вашем traffic.

Quantization уменьшает model, сохраняя ту же function в меньшем числе bits. Distillation уменьшает его, обучая smaller model имитировать larger one11 — идея, появившаяся почти за decade до deep learning.12

Тонкость в том, на чём учится student. Не на правильном ответе: на нём его можно было обучить напрямую. То, что добавляет teacher, — whole distribution. Спросите model, что следует за phrase, и посмотрите дальше argmax:

TEXT
"She poured the milk into the"
  ' jug' 0.1355   ' cup' 0.1051   ' bowl' 0.0605   ' large' 0.0380   ' milk' 0.0360

Hard label говорит jug и больше ничего. Soft label говорит jug, а также что cup было почти так же хорошо, bowl plausible, и large — adjective, совершенно другое grammatical continuation — всё ещё живо. Это исходный аргумент: это 7, но выглядит довольно похоже на 1, и resemblance — information, которую hard label выбрасывает.

Поэтому distillation использует temperature. Деление logits на TT перед softmax выравнивает distribution и повышает relative weight runners-up: на этой phrase ratio между top token и third падает с 2,24 при T=1T = 1 до 1,50 при T=2T = 2 — square root of the first, что и делает деление logits на два с ratio. Тот же ordering, больше attention loss на near misses. Gradient student несёт uncertainty teacher, а не только его verdict.

Всё в этой главе теперь — одна сумма:

memory=N×bytes per weightfixed+T×2LHkvdhead×bytesgrows with every token+runtime overheadcall it 1.5 GB\text{memory} = \underbrace{N \times \text{bytes per weight}}_{\text{fixed}} + \underbrace{T \times 2 L H_{kv} d_{\text{head}} \times \text{bytes}}_{\text{grows with every token}} + \underbrace{\text{runtime overhead}}_{\text{call it 1.5 GB}}

где TTобщее число resident token по всем concurrent requests. Применим её: строки 7B и 70B предполагают 8 key-value heads размерности 128, строка 13B — полный multi-head attention с 40 heads, как были устроены те поколения model, — и это заметно.

8 GB

modelprecisionweightsfree after overheadcontext tokens that fit
7Bfp1613,0 GBdoes not fit
7Bint86,5 GBdoes not fit
7Bint4 (g128)3,4 GB3,1 GB25 710
13Bint4 (g128)6,2 GB0,3 GB337
70Bint4 (g128)33,6 GBdoes not fit

16 GB

modelprecisionweightsfree after overheadcontext tokens that fit
7Bfp1613,0 GB1,5 GB11 972
7Bint86,5 GB8,0 GB65 378
7Bint4 (g128)3,4 GB11,1 GB91 246
13Bint812,1 GB2,4 GB3 136
13Bint4 (g128)6,2 GB8,3 GB10 822

24 GB

modelprecisionweightsfree after overheadcontext tokens that fit
7Bfp1613,0 GB9,5 GB77 508
7Bint86,5 GB16,0 GB130 914
7Bint4 (g128)3,4 GB19,1 GB156 782
13Bint812,1 GB10,4 GB13 622
13Bint4 (g128)6,2 GB16,3 GB21 308
70Bint4 (g128)33,6 GBdoes not fit

Посмотрите на строку 13B в таблице 8 GB. Weights помещаются — 6,2 GB из 8, — поэтому в обычной манере речи 13B model «работает на карте 8 GB». У него 337 token context, что не conversation, а едва prompt. «Помещается ли он» — неправильный вопрос. Правильный: «с каким context и для скольких пользователей одновременно».

Посмотрите также на две строки int8 в 16 GB. 7B получает 65 378 token, а 13B — 3 136: двадцатикратная разница из-за 5,6 GB дополнительных weights, потому что у этого 13B multi-head attention и его cache стоит 800 KB на token против 128 KB у 7B. Два model похожего size, один unusable для long context, по причине, которая не появляется в headline ни одной model card.

Тринадцать глав назад это был perceptron с двумя weights и bias. Теперь это transformer, который был designed, trained, aligned, научен тратить compute на сложные вопросы и served с измеренной стоимостью на token — и в нём не осталось ни одного unopened box.

Здесь это заканчивается, и заканчивается намеренно.

Глава 14 начинается с model, находящегося где-то ещё. Не в вашем process, не в вашей memory, не в variable, который можно print: на машине, которую вы не администрируете, за API key, port и bill. Всё измеренное здесь всё ещё происходит — prefill всё ещё выполняется перед first token, cache всё ещё растёт вместе с conversation, batch, в котором вы находитесь, всё ещё принадлежит кому-то другому и всё ещё решает вашу latency, — но теперь вы наблюдаете это через stream of Server-Sent Events, finish_reason и HTTP 429 с header Retry-After. Вопросы меняются вместе с vantage point: не как вычисляется этот gradient, а почему мой счёт утроился. Меняется и language, и глава 14 объясняет это правило, а не объявляет его: до этого code держал weights, gradients, logits и tokenizer bytes; дальше он держит connection, retry, cancellation и accumulated state. Тринадцать глав позади не отбрасываются при переходе. Они — описание того, что работает по другую сторону port.


Два omission намеренны. FlashAttention (Dao et al., arXiv:2205.14135) — не другой attention: он вычисляет ту же function, разбивая operation на tiles так, чтобы score matrix n×nn \times n никогда не записывалась в memory, поэтому 67 MB во второй таблице этой главы на практике меньше, чем suggests arithmetic. А сами kernels delegated: lecture 10 Stanford CS336 покрывает inference systems в глубине, на которую эта глава не претендует, а repository llama.cpp и specification GGUF — primary sources для CPU side.

  1. Shazeer, N. Fast Transformer Decoding: One Write-Head is All You Need. arXiv:1911.02150 (2019). Статья в значительной степени является аргументом о memory-bandwidth, и так и читается.

  2. Ainslie, J. et al. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. arXiv:2305.13245 (2023). Включает uptraining recipe, который converts existing multi-head checkpoint, поэтому GQA распространился так быстро.

  3. Yu, G.-I., Jeong, J. S., Kim, G.-W., Kim, S. and Chun, B.-G. Orca: A Distributed Serving System for Transformer-Based Generative Models. OSDI 2022. Представляет iteration-level scheduling — continuous batching — и selective batching.

  4. Kwon, W. et al. Efficient Memory Management for Large Language Model Serving with PagedAttention. arXiv:2309.06180 (2023), SOSP 2023. Статья, на которой построен vLLM; §3 полностью раскрывает аналогию с operating systems.

  5. Dettmers, T., Pagnoni, A., Holtzman, A. and Zettlemoyer, L. QLoRA: Efficient Finetuning of Quantized LLMs. arXiv:2305.14314 (2023). NF4 определён в §3; sixteen level values, использованные в измерении выше, — те, что выводит эта статья.

  6. Dettmers, T., Lewis, M., Belkada, Y. and Zettlemoyer, L. LLM.int8(): 8-bit Matrix Multiplication for Transformers at Scale. arXiv:2208.07339 (2022). Анализ outlier-feature в §4 — источник phenomenon, измеренного выше, включая вывод, что outliers систематически появляются at scale. 2

  7. Xiao, G., Lin, J., Seznec, M., Wu, H., Demouth, J. and Han, S. SmoothQuant: Accurate and Efficient Post-Training Quantization for Large Language Models. arXiv:2211.10438 (2022).

  8. Frantar, E., Ashkboos, S., Hoefler, T. and Alistarh, D. GPTQ: Accurate Post-Training Quantization for Generative Pre-trained Transformers. arXiv:2210.17323 (2022).

  9. Lin, J. et al. AWQ: Activation-aware Weight Quantization for LLM Compression and Acceleration. arXiv:2306.00978 (2023).

  10. Leviathan, Y., Kalman, M. and Matias, Y. Fast Inference from Transformers via Speculative Decoding. arXiv:2211.17192 (2022). Theorem 1 — proof, что output distribution не меняется; Chen et al. (arXiv:2302.01318) опубликовали ту же идею независимо.

  11. Hinton, G., Vinyals, O. and Dean, J. Distilling the Knowledge in a Neural Network. arXiv:1503.02531 (2015). Temperature и аргумент о «dark knowledge».

  12. Buciluă, C., Caruana, R. and Niculescu-Mizil, A. Model Compression. KDD 2006. Distillation на девять лет раньше, для ensembles, а не transformers.

Готовы доверить выбор модели LIA?

Создавайте со всеми ИИ-моделями в одном месте — начните бесплатно уже сегодня.