Serving · Memory
- model
- GPT-2 small 124M weights
- fp32
- 498 MB
- int8
- 124 MB
- int4
- 62 MB + scales
- perplexity
- 52.9 → 62.6 fp32 → int4 g128
All steps
Fewer bits per number
Models train in 32- or 16-bit floating point. Serving them in 8 or 4 bits makes the weights 2 to 4 times smaller, so a model fits on fewer GPUs and each decode step, which is limited by reading the weights, runs faster.
32 → 16 → 8 → 4 bits per weightRounding onto a grid
The simplest scheme, absmax: divide by a scale so the largest weight lands on the largest integer, round, and multiply back when computing. int8 has 255 levels, int4 only 15, so the rounding error is much larger. A real column of GPT-2’s weights.
q = round(w / s) · s = max|w| / (2^(b−1) − 1)One scale or many
One scale for a whole matrix is set by its single largest weight, which makes the grid coarse for everything else. A scale per output channel, or per group of 128 weights, follows the local range. Real GPT-2, all linear layers quantized and rerun.
per tensor · per channel · per group of 128Outliers
Activations are harder than weights: a few dimensions are far larger than the rest. In this real GPT-2 activation one number is 30 times the typical size, and with one int4 scale most values round to zero. LLM.int8 keeps such dimensions in 16 bits; most 4-bit serving quantizes only the weights.
one large value stretches the whole gridWhat it costs the model
The real cost for GPT-2 small. int8 with a scale per channel is nearly free; int4 in groups of 128 costs a little; int4 with a single scale per matrix breaks the model. Methods like GPTQ and AWQ choose the rounding more cleverly and keep 4-bit models close to full precision.
perplexity on a paragraph · next-token guesses
Code
# absmax, symmetric: one scale per group of weights
q = 2 ** (bits - 1) - 1 # 127 for int8, 7 for int4
W = W.reshape(out_features, -1, group) # groups of 128 along the input
scale = W.abs().amax(dim=-1, keepdim=True) / q
codes = torch.round(W / scale).clamp(-q, q).to(torch.int8)
# at inference: dequantize on the fly (weight-only) and multiply in 16-bit
y = x @ (codes * scale).reshape(out_features, -1).T