PidokuInfra

Project 10 — Quantized Inference

Advanced 10h Difficulty 4/5 Topic 10 of 15

Prerequisites III.11-12, Section VII (02-06), Project 05

Shrink the weights, measure what you lost and what you gained, and learn why the two are not the same question.


1. What you build#

Two things:

  1. A quantizer from scratch — round-to-nearest INT8 and INT4, per-tensor vs per-channel vs group-wise, symmetric vs asymmetric — applied to GPT-2, with perplexity measured for each.
  2. A production-style pipeline — GPTQ or AWQ via an existing library on a ≤1B modern model, served, and benchmarked for memory, tokens/s, and quality against the FP16 baseline.

Diagram — The pipeline#

flowchart LR
  W["FP16 weights"] --> QZ["Quantize<br/>RTN / GPTQ / AWQ"]
  CAL["Calibration data"] --> QZ
  QZ --> EV["Evaluate<br/>perplexity + task outputs"]
  QZ --> BN["Benchmark<br/>bytes, tokens/s, memory"]
  EV --> D{"Quality within budget?"}
  BN --> D
  D -->|"yes"| SHIP["Ship"]
  D -->|"no"| MX["More bits, per-channel,<br/>or mixed precision"]
  MX --> QZ

  class W,CAL neutral
  class QZ memory
  class EV,BN compute
  class D queue
  class SHIP io
  class MX warn

2. Why it matters#

Decode is bandwidth-bound: every token reads every weight. Halving the bytes roughly halves the time and the memory — if the kernel can consume the compressed format directly. Whether it can is the whole story, and it is why “I quantized it and it got slower” is one of the most common surprises in the field.


3. Read first#


4. Spec#

Part A — from scratch (GPT-2 small, weight-only)
  quantize(W, bits, scheme)  -> (q, scale, zero_point)
  dequantize(q, scale, zp)   -> W_hat
  schemes: per-tensor | per-channel (per output row) | group-wise (g = 64 or 128)
           symmetric | asymmetric
  layers:  all Linear weights; keep embeddings, LayerNorm, and biases in float
  metric:  perplexity on WikiText-2 test (stride 512), plus per-layer relative MSE

Part B — library pipeline (≤1B model, e.g. Qwen2.5-0.5B)
  FP16 baseline | INT8 | INT4-RTN | INT4-GPTQ or AWQ   (llm-compressor / AutoAWQ / GGUF)
  metrics: weight bytes, peak memory, tokens/s at batch 1 and batch 16, perplexity,
           and exact-match on 50 prompts vs FP16 greedy output

5. Milestones#

  1. Baseline perplexity for FP32 GPT-2. Get this right first; every later number is relative to it.
  2. RTN INT8, per-tensor. Quantize → dequantize → evaluate. (“Fake quant”: you measure accuracy, not speed.)
  3. Per-channel. Same bits, better quality. Look at one weight matrix’s row ranges to see why.
  4. INT4. Per-tensor falls apart; group-wise recovers most of it. Plot perplexity vs effective bits per weight, including the scale/zero-point overhead.
  5. Sensitivity scan. Quantize one layer at a time to INT4 and record the perplexity delta. Which layers are fragile?
  6. Real kernels. Run Part B. Now the weights stay compressed in memory and the runtime dequantizes on the fly.
  7. The report. One table: bytes, tokens/s, perplexity, for each variant.

6. Starter skeleton#

Python
def quantize_sym(W, bits, dim=None):
    """Symmetric. dim=None → per-tensor; dim=1 → one scale per output row."""
    qmax = 2 ** (bits - 1) - 1
    amax = W.abs().amax(dim=dim, keepdim=dim is not None).clamp(min=1e-8)
    scale = amax / qmax
    q = torch.clamp(torch.round(W / scale), -qmax - 1, qmax)
    return q.to(torch.int8), scale

def quantize_groupwise(W, bits, g=128):
    out, inp = W.shape
    Wg = W.reshape(out, inp // g, g)
    qmax = 2 ** (bits - 1) - 1
    scale = Wg.abs().amax(-1, keepdim=True).clamp(min=1e-8) / qmax
    q = torch.clamp(torch.round(Wg / scale), -qmax - 1, qmax)
    return (q * scale).reshape(out, inp), scale       # fake-quant result

def effective_bits(bits, g, scale_bits=16):
    return bits + scale_bits / g                      # INT4, g=128 → 4.125 bits/weight

@torch.inference_mode()
def perplexity(model, ids, ctx=1024, stride=512):
    nll, n = 0.0, 0
    for i in range(0, ids.size(1) - 1, stride):
        x = ids[:, max(0, i + stride - ctx): i + stride]
        tgt = x.clone(); tgt[:, :-stride] = -100      # score only the new tokens
        nll += model(x, labels=tgt).loss.item() * stride; n += stride
    return math.exp(nll / n)

7. What to measure#

MeasurementExpectation to write down first
Perplexity: FP32 → INT8 per-tensor → INT8 per-channelNear-lossless at INT8
Perplexity: INT4 per-tensor vs group-wise vs GPTQ/AWQWide spread; calibration matters at 4 bits
Effective bits/weight including metadataNot exactly 4
Per-layer sensitivityA few layers dominate the error
Weight bytes and peak memory per variant~½ and ~¼ of FP16
Tokens/s at batch 1, per variantSpeedup if the kernel reads compressed weights
Tokens/s at batch 16+Gain shrinks — you are leaving the memory-bound regime
Fake-quant speed vs FP16Identical or slower — nothing got smaller in memory

8. Done when#

  • You reproduced “INT8 is nearly free, naive INT4 is not” with your own numbers.
  • You have the per-layer sensitivity chart.
  • Part B table is complete and you can explain every speedup and every non-speedup.
  • You can state why weight-only INT4 helps decode more than prefill (Checkpoint D).
  • You checked outputs on real prompts, not perplexity alone.

9. Common pitfalls#

Benchmarking fake quantization for speed. The tensor is still float in memory.

Quantizing embeddings and norms for negligible savings and real damage.

Calibrating on data that does not look like production traffic.

Trusting perplexity alone. Tool calling, long-context recall, and structured output degrade before perplexity moves. Test your actual task.

Ignoring metadata overhead when quoting bits per weight.

Comparing runtimes instead of formats. A GGUF build on llama.cpp vs an FP16 PyTorch loop differs in far more than quantization.


10. Stretch goals#

  • Implement GPTQ’s core update for one layer yourself (Hessian from calibration activations, column-by-column error compensation) and beat your RTN result.
  • FP8 weights + activations on a GPU that supports it; compare to INT8.
  • Quantize the KV cache to INT8/FP8 and measure concurrency gained vs quality lost (VII.12).
  • Mixed precision from your sensitivity scan: fragile layers at 8 bits, the rest at 4.

11. Interview questions this project answers#

  1. Why is per-channel quantization better than per-tensor at equal bit width?
  2. What do GPTQ and AWQ do that round-to-nearest does not?
  3. Why can a quantized model be slower?
  4. Why does quantization help decode more than prefill?
  5. How do you decide whether a quantized model is safe to ship?

12. Next#

Project 11 — GPU benchmark suite

↑↓ navigate↵ openesc close