Serving · Memory

Quantization

weights in 8 and 4 bits, on the real GPT-2

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

  1. 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 weight

  2. Rounding 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)

  3. 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 128

  4. Outliers

    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 grid

  5. What 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

Go deeper