Trained networks contain astonishing redundancy. Han et al. (2015) showed that most weights of AlexNet and VGG could be removed with little loss in accuracy. Compressing models matters wherever resources are limited: smartphones, microcontrollers in the field, browsers, or servers handling millions of requests where every millisecond and watt counts. The two main tools are pruning โ removing parameters โ and quantisation โ storing and computing with fewer bits.
Pruning#
Unstructured pruning#
Remove individual weights, most commonly those with the smallest magnitude (magnitude pruning). The resulting sparse matrices can reach very high sparsity (80โ95% on many networks) with small accuracy loss when combined with fine-tuning. But unstructured sparsity rarely speeds up standard hardware, which is optimised for dense matrix multiplication; it mainly reduces storage (with sparse formats) unless special kernels or hardware are used.
Structured pruning#
Remove entire channels, filters, attention heads or layers. The model becomes genuinely smaller and faster on ordinary hardware, at the cost of lower achievable sparsity. Importance can be measured by the norm of a filter's weights, the scaling factors of BatchNorm layers (network slimming), or the effect on the loss (Taylor-expansion criteria).
A middle ground, N:M semi-structured sparsity (e.g. 2 non-zeros in every block of 4), is accelerated by recent GPUs.
The pruning workflow#
- Train a dense model.
- Prune a fraction of weights by some importance criterion.
- Fine-tune to recover accuracy.
- Repeat (iterative pruning usually beats one-shot pruning at high sparsity).
import torch
import torch.nn as nn
import torch.nn.utils.prune as prune
model = nn.Sequential(nn.Linear(784, 300), nn.ReLU(), nn.Linear(300, 100), nn.ReLU(), nn.Linear(100, 10))
params = [(m, "weight") for m in model if isinstance(m, nn.Linear)]
prune.global_unstructured(params, pruning_method=prune.L1Unstructured, amount=0.8) # prune 80% globally
total = sum(m.weight.nelement() for m, _ in params)
zeros = sum((m.weight == 0).sum().item() for m, _ in params)
print(f"global sparsity: {zeros / total:.1%}")
for m, _ in params:
prune.remove(m, "weight") # make pruning permanent (bake mask into weights)The lottery ticket hypothesis#
Frankle and Carbin (2019) found that dense networks contain small sub-networks ("winning tickets") that, when reset to their original initialisation and trained in isolation, match the full network's accuracy. This suggests over-parameterisation helps optimisation by providing many candidate sub-networks, rather than being needed for the final function. For large networks, rewinding to an early training checkpoint rather than initialisation works better.
Quantisation#
Replace 32-bit floating-point numbers with low-bit representations, typically 8-bit integers (INT8) or even 4-bit. Benefits: 4โ8ร smaller models, less memory bandwidth, and faster integer arithmetic on CPUs, mobile NPUs and GPUs.
Affine (asymmetric) quantisation maps a real value $x$ to an integer $q$ with a scale $s$ and zero point $z$:
For $b$ bits and a real range $[x_{\min}, x_{\max}]$: $s = \frac{x_{\max} - x_{\min}}{2^b - 1}$. Symmetric quantisation fixes $z = 0$.
import numpy as np
def quantize(x, bits=8):
qmin, qmax = 0, 2**bits - 1
s = (x.max() - x.min()) / (qmax - qmin)
z = np.round(qmin - x.min() / s)
q = np.clip(np.round(x / s + z), qmin, qmax).astype(np.int32)
return q, s, z
def dequantize(q, s, z):
return s * (q - z)
w = np.random.default_rng(0).normal(0, 0.05, 10_000)
for bits in [8, 4, 2]:
q, s, z = quantize(w, bits)
err = np.abs(dequantize(q, s, z) - w).mean()
print(f"{bits}-bit: mean abs error {err:.5f} ({err / np.abs(w).mean():.1%} of mean |w|)")Granularity#
- Per-tensor: one scale for the whole tensor โ simple but hurt by outliers.
- Per-channel: one scale per output channel โ standard for weights.
- Per-group: one scale per block of (e.g.) 64โ128 weights โ standard for 4-bit LLM weights.
Post-training quantisation vs quantisation-aware training#
- Post-training quantisation (PTQ): quantise a trained model directly. Dynamic PTQ quantises weights ahead of time and activations on the fly; static PTQ uses a small calibration set to fix activation ranges. Fast and often sufficient for INT8.
- Quantisation-aware training (QAT): simulate quantisation during training ("fake quantisation") and backpropagate with the straight-through estimator (treat rounding as identity in the backward pass). Recovers accuracy at low bit-widths.
Quantising large language models#
LLM activations contain rare, very large outlier features that break naive INT8 quantisation. Methods such as LLM.int8() (handles outliers in higher precision), SmoothQuant (migrates difficulty from activations to weights), GPTQ and AWQ (accurate 4-bit weight-only quantisation using second-order or activation-aware criteria) make it possible to run large models on consumer GPUs and laptops with modest quality loss. Formats such as GGUF package quantised models for local inference.
Choosing a compression strategy#
| Goal | Recommended |
|---|---|
| Faster CPU/mobile inference | INT8 static PTQ; structured pruning; distillation to a small architecture |
| Fit an LLM on a small GPU | 4-bit weight-only quantisation (GPTQ/AWQ) |
| Microcontroller (TinyML) | Small architecture + INT8 QAT |
| Maximum accuracy at low bits | QAT, per-channel/group scales |