Course Content
Model Deployment for AI Engineers
4 sections · 10 lessons
Quantization and Pruning
A team ships a BERT-base sentiment classifier to a mobile app. Offline it scores 92.1% F1. The artefact is 440 MB. The app store rejects the build — the limit for over-the-air download is 200 MB — and the product manager asks the obvious question: can it be smaller?
The instinctive answer is "retrain something smaller". That costs a week of GPU time and usually loses several points of accuracy. The better answer is that the 440 MB is 110 million numbers stored at 32 bits each, and almost none of them need 32 bits. Storing them at 8 bits gives 110 MB, fits the limit, runs about 2.5× faster on the phone's CPU, and — done properly — loses about 0.4 points of F1.
That trade is the entire subject of this lesson. Making a model smaller and faster is not a matter of "using less"; it is a set of specific, measurable techniques with specific, measurable costs. Getting them right is often the difference between a model that ships and a model that does not.
Why this is a separate discipline
Training optimises for accuracy on a validation set. Serving optimises for something else entirely: latency at the 99th percentile, throughput per pound of compute, memory footprint, and energy per inference. A model that is excellent by the first metric can be unusable by the others.
| Constraint | Typical requirement | What violating it looks like |
|---|---|---|
| Mobile app bundle | Under ~150 MB | Store rejection, or users abandoning the download |
| Interactive API | p99 under 200 ms | Users perceive the product as broken |
| Edge device RAM | 512 MB total, shared | Out-of-memory kill mid-inference |
| Cost per million inferences | Set by unit economics | The feature is unprofitable and gets cancelled |
| Battery / thermal | Sustained load without throttling | Device heats, clocks drop, latency doubles |
Four techniques dominate. Quantization stores and computes with fewer bits per number. Pruning removes weights or whole channels. Mixed precision runs parts of the computation in 16-bit. Distillation trains a small model to imitate a large one. They attack different parts of the problem and combine well.
Quantization: fewer bits per number
What precision actually costs you
A 32-bit float uses 1 sign bit, 8 exponent bits, and 23 mantissa bits — roughly seven decimal digits of precision across a range up to about 3.4×1038. Neural network weights, after training, essentially never use that range. A typical trained layer has weights concentrated in something like [-0.3, 0.3], with occasional outliers. You are spending 32 bits to store a number that lives in a narrow band.
| Type | Bits | Range | 110M-param model | Relative memory bandwidth |
|---|---|---|---|---|
| FP32 | 32 | ±3.4e38 | 440 MB | 1.0× |
| FP16 | 16 | ±65,504 | 220 MB | 0.5× |
| BF16 | 16 | ±3.4e38 | 220 MB | 0.5× |
| INT8 | 8 | -128 to 127 | 110 MB | 0.25× |
| INT4 | 4 | -8 to 7 | 55 MB | 0.125× |
Memory bandwidth is the column that explains the speed-up. For most inference workloads, the bottleneck is not arithmetic — it is moving weights from RAM into the CPU or GPU. Quarter the bytes and you roughly quarter the time spent waiting on memory, before any arithmetic speed-up from integer units.
The mapping, with real arithmetic
Quantization maps a float range onto the integers. Affine (asymmetric) quantization uses a scale s and a zero-point z:
with
Take a real tensor whose observed range is [-2.5, 3.1], targeting signed INT8 (range -128 to 127).
Now quantize the value 1.7:
- q=round(1.7/0.021961)+(−14)=round(77.41)−14=77−14=63
- Dequantize: x^=0.021961×(63−(−14))=0.021961×77=1.6910
- Error: ∣1.7−1.6910∣=0.0090
The error is bounded by half the scale, s/2=0.01098, and 0.0090 sits under it. Check the endpoints: x=3.1 gives q=round(141.16)−14=127, exactly the top of the integer range, and x=−2.5 gives q=−114−14=−128, exactly the bottom. The mapping uses all 256 available levels.
Why one outlier ruins everything
This is the failure mode that catches people, and the arithmetic makes it obvious. Suppose one weight in that tensor is 20.0 instead of 3.1 — a single outlier out of a million values. The range becomes [-2.5, 20.0]:
The maximum error jumps from 0.011 to 0.044 — four times worse — for every one of the other 999,999 weights, all so that one outlier can be represented exactly. The 999,999 useful weights now occupy only about 63 of the 256 available integer levels; the rest are reserved for a gap containing nothing.
Quantization error is set by the widest value in the tensor, so a single outlier degrades the precision of every other weight it shares a scale with.
Two standard defences. Per-channel quantization computes a separate scale for each output channel rather than one for the whole tensor, so an outlier only damages its own channel — this is nearly free and should be the default for convolution and linear weights. Percentile calibration clips the range at, say, the 99.99th percentile instead of the true maximum, accepting clipping error on a handful of values in exchange for much finer resolution everywhere else.
Symmetric versus asymmetric
Symmetric quantization forces z=0 and uses s=max(∣xmin∣,∣xmax∣)/127. For our original tensor that gives s=3.1/127=0.024409 — slightly coarser than the affine scale of 0.021961, because the range [-3.1, -2.5] is wasted.
The compensation is speed. With z=0, a quantized matrix multiply is a plain integer dot product; with a non-zero zero-point, expanding (q1−z1)(q2−z2) produces cross terms that need extra work. In practice: symmetric for weights (roughly zero-centred anyway, and speed matters in the inner loop), asymmetric for activations (post-ReLU activations are all non-negative, so symmetric would throw away half the integer range).
Post-training quantization
PTQ takes a trained FP32 model and converts it with no retraining. Two variants follow. The code uses PyTorch's eager-mode torch.ao.quantization API because it shows each step plainly; current PyTorch marks that module as deprecated and points to the separate torchao library (its quantize_ API, and the prepare_pt2e / convert_pt2e flow for static quantization). The concepts — observers, calibration, fusion, fake-quant — carry over unchanged. For a model you are exporting anyway, quantizing the ONNX file with ONNX Runtime, as in the first lesson, is often the simplest route.
1import torch2from torch.ao.quantization import quantize_dynamic34# Dynamic: weights quantized ahead of time, activations quantized on the fly.5model_int8 = quantize_dynamic(6 model, {torch.nn.Linear, torch.nn.LSTM}, dtype=torch.qint87)8torch.save(model_int8.state_dict(), "model_int8.pt")Dynamic quantization needs no calibration data at all, because activation ranges are measured per batch at run time. That measurement costs something, so the speed-up is smaller — but it is the right first thing to try for transformers and RNNs, where linear layers dominate.
Static quantization also fixes the activation scales ahead of time, which requires showing the model representative data:
1import torch.ao.quantization as tq23model.eval()4model.qconfig = tq.get_default_qconfig("x86") # per-channel weights, histogram observer56model_fused = tq.fuse_modules(model, [["conv1", "bn1", "relu1"]])7model_prepared = tq.prepare(model_fused)89# Calibration: 100-500 REAL samples. No labels needed, no gradients.10with torch.no_grad():11 for batch, _ in calibration_loader:12 model_prepared(batch)1314model_int8 = tq.convert(model_prepared)Two details decide whether this works.
Fusion first. fuse_modules folds batch-norm into the preceding convolution and merges the activation. This is not merely an optimisation — an unfused conv-then-batchnorm sequence quantizes the intermediate tensor between them, and that intermediate often has a much wider range than either the input or the output. Skipping fusion is one of the most common causes of a surprisingly large accuracy drop.
Calibration data must be real and representative. Random noise gives activation ranges that have nothing to do with production inputs. If your training set is 60% daytime photos and production is 90% night-time, calibrate on night-time photos. 100 to 500 samples is plenty; 10,000 adds nothing.
Quantization-aware training
When PTQ costs too much accuracy, QAT simulates quantization during fine-tuning so the weights adapt to it. Fake-quant nodes round activations and weights in the forward pass; gradients flow through using a straight-through estimator, since round has zero derivative almost everywhere.
1model.train()2model.qconfig = tq.get_default_qat_qconfig("x86")3model_qat = tq.prepare_qat(tq.fuse_modules(model, [["conv1", "bn1", "relu1"]]))45optimiser = torch.optim.SGD(model_qat.parameters(), lr=1e-4) # 10-100x below original6for epoch in range(3): # a few epochs, not a full run7 for x, y in train_loader:8 loss = criterion(model_qat(x), y)9 optimiser.zero_grad(); loss.backward(); optimiser.step()1011model_qat.eval()12model_int8 = tq.convert(model_qat)The low learning rate matters. You are nudging an already-good solution toward one that survives rounding, not re-training. A normal learning rate destroys the pretrained weights and you end up worse than PTQ.
| Dynamic PTQ | Static PTQ | QAT | |
|---|---|---|---|
| Data needed | None | 100–500 unlabelled samples | Full labelled training set |
| Time to apply | Seconds | Minutes | Hours to days |
| Typical accuracy drop | 0.5–2% | 1–3% | 0.1–0.5% |
| Typical CPU speed-up | 1.5–2× | 2–4× | 2–4× |
| Best for | Transformers, RNNs | CNNs | When PTQ loses too much |
Always start with static PTQ. Measure. Only escalate to QAT if the measured drop exceeds your budget, because QAT costs a training cycle.
Mixed precision: 16-bit where it is safe
Mixed precision keeps a master copy of weights in FP32 but runs most operations in 16-bit. On GPUs with tensor cores this roughly doubles arithmetic throughput and halves activation memory.
1from torch.amp import autocast, GradScaler23scaler = GradScaler("cuda")4for x, y in loader:5 optimiser.zero_grad()6 with autocast("cuda", dtype=torch.float16):7 loss = criterion(model(x), y) # matmuls in FP16, softmax/norms in FP328 scaler.scale(loss).backward() # scale up before backward9 scaler.step(optimiser) # unscale, then step (skips if inf/nan)10 scaler.update()autocast maintains a list of which operations are safe in FP16. Matrix multiplications and convolutions are; reductions, softmax, and normalisation layers stay in FP32 because summing many small numbers in FP16 loses them.
Why the gradient scaler exists
FP16's smallest normal positive value is about 6.1×10−5. Gradients in a deep network routinely reach 10−8. In FP16 that is exactly zero — the weight simply stops learning, silently, and your loss curve flattens for no visible reason.
GradScaler multiplies the loss by a large factor (starting at 65,536 = 216) before backward(). By the chain rule every gradient is scaled by the same factor, so 1×10−8 becomes 6.55×10−4 — comfortably representable. Before the optimiser step, gradients are divided back down. If any gradient overflowed to infinity, the step is skipped and the scale factor halved; after many clean steps the scale is doubled again.
Mixed-precision training without loss scaling does not crash — it quietly zeroes your smallest gradients, which is much harder to diagnose than a crash.
BF16
BFloat16 splits its 16 bits differently: 8 exponent bits (identical to FP32) and 7 mantissa bits. It has FP32's dynamic range and about one-eighth of FP16's precision (7 mantissa bits against 10).
| FP16 | BF16 | |
|---|---|---|
| Exponent / mantissa bits | 5 / 10 | 8 / 7 |
| Max magnitude | 65,504 | ~3.4e38 |
| Smallest normal | 6.1e-5 | ~1.2e-38 |
| Needs loss scaling | Yes | No |
| Hardware | Volta onward, most GPUs | Ampere onward, TPUs, recent Xeon |
If your hardware supports BF16, prefer it for training: no scaler, no overflow tuning, no silent underflow. FP16 retains an edge for inference on older GPUs and where the extra mantissa bits genuinely help.
Pruning: removing weights
The observation behind it
Trained networks are heavily over-parameterised. In a typical ResNet-50, 30–50% of weights have magnitude near zero and contribute almost nothing to any output. Pruning removes them.
1import torch.nn.utils.prune as prune23prune.l1_unstructured(model.fc1, name="weight", amount=0.5) # zero the smallest 50%4print((model.fc1.weight == 0).float().mean().item()) # 0.556prune.remove(model.fc1, "weight") # bake the mask in permanentlyNow the crucial disappointment. That layer has 50% zeros — and it is exactly the same size and exactly the same speed. The zeros are still stored as FP32 zeros, and a dense matrix multiply still multiplies by them. Unstructured pruning only pays off with a sparse storage format and hardware or kernels that exploit it, and general-purpose CPUs and GPUs mostly do not, below about 90% sparsity.
Structured pruning, which actually helps
Structured pruning removes whole channels or filters, so the tensor genuinely shrinks and every framework gets faster for free.
prune.ln_structured(model.conv2, name="weight", amount=0.3, n=2, dim=0)# dim=0 = output channels; ranks channels by L2 norm, removes the weakest 30%Work out what that saves. A 3×3 convolution with 512 input and 512 output channels holds
Prune 30% of channels throughout the network, so this layer sees 358 inputs and produces 358 outputs:
That is a 51.1% reduction, not 30% — because the channel cut applies to both dimensions and they multiply. FLOPs fall by the same proportion. This compounding is why structured pruning at a modest per-layer rate produces large end-to-end savings.
| Unstructured | Structured | |
|---|---|---|
| What is removed | Individual weights | Whole channels or filters |
| Accuracy at equal parameter count | Better | Worse |
| Real speed-up on standard hardware | Essentially none below ~90% sparsity | Directly proportional |
| Disk size after compression | Smaller (zeros compress well) | Smaller (fewer values) |
| Use when | Targeting sparse accelerators or download size | You need actual latency reduction |
Prune iteratively, not in one shot
Removing 70% of weights at once typically drops accuracy by 15–20 points, and fine-tuning does not fully recover it. Removing 20% at a time with a short fine-tune between rounds reaches the same sparsity with a fraction of the damage — the surviving weights get a chance to compensate for each removal before the next.
1target, per_round = 0.70, 0.202current = 0.03while current < target:4 step = min(per_round, (target - current) / (1 - current))5 for module in prunable_layers(model):6 prune.l1_unstructured(module, name="weight", amount=step)7 fine_tune(model, epochs=2, lr=1e-4)8 current = 1 - (1 - current) * (1 - step)9 print(f"sparsity {current:.2%} acc {evaluate(model):.4f}")Note the compounding: pruning 20% of what remains, three times, gives 1−0.83=48.8% sparsity, not 60%. The formula in the loop accounts for that so you actually land on 70%.
Knowledge distillation: training a small model to imitate a large one
Why not just train the small model directly on the labels? Because a hard label carries one bit of relevant information — "this is a cat". The large model's full output distribution carries much more: it says this image is 91% cat, 6% lynx, 2% dog, 0.1% aeroplane. That ranking encodes a similarity structure the labels never mention, and a small model learns far faster from it. Hinton called it "dark knowledge".
The catch is that a confident teacher's probabilities are nearly one-hot, so the extra information is numerically invisible. Temperature fixes that by dividing logits before the softmax. Take teacher logits [8.0, 2.0, 1.0]:
| Temperature | Class A | Class B | Class C |
|---|---|---|---|
| T=1 | 0.9966 | 0.00247 | 0.00091 |
| T=4 | 0.7159 | 0.1597 | 0.1244 |
At T=1 the runner-up carries a gradient of essentially zero. At T=4 — logits become [2.0, 0.5, 0.25], giving e2.0=7.389, e0.5=1.649, e0.25=1.284 over a sum of 10.322 — the fact that B outranks C is a real, learnable signal.
1import torch.nn.functional as F23def distillation_loss(student_logits, teacher_logits, labels, T=4.0, alpha=0.7):4 soft = F.kl_div(5 F.log_softmax(student_logits / T, dim=1),6 F.softmax(teacher_logits / T, dim=1),7 reduction="batchmean",8 ) * (T * T) # see note below9 hard = F.cross_entropy(student_logits, labels)10 return alpha * soft + (1 - alpha) * hard1112teacher.eval()13for x, y in loader:14 with torch.no_grad():15 t_logits = teacher(x)16 loss = distillation_loss(student(x), t_logits, y)17 optimiser.zero_grad(); loss.backward(); optimiser.step()The T2 factor is not decoration. Dividing logits by T scales the gradients of the soft loss by roughly 1/T2, so without the correction the soft term would contribute 16× less at T=4 than at T=1, and changing the temperature would silently change the effective balance between the two loss terms.
Real results from this technique: DistilBERT is 40% smaller and 60% faster than BERT-base while retaining about 97% of its GLUE score. That is a better trade than you get from training a 6-layer transformer from scratch on the same data.
Choosing and combining
| Technique | Size reduction | Typical speed-up | Accuracy cost | Effort |
|---|---|---|---|---|
| Dynamic INT8 PTQ | 4× | 1.5–2× | 0.5–2% | Minutes |
| Static INT8 PTQ | 4× | 2–4× | 1–3% | Hours |
| QAT | 4× | 2–4× | 0.1–0.5% | Days |
| FP16 / BF16 | 2× | 1.5–3× on tensor cores | ~0% | Hours |
| Structured pruning 30% | ~1.5–2× | ~1.4–1.9× | 1–3% | Days |
| Distillation | 2–10× | 2–10× | 2–5% | Weeks |
They stack multiplicatively, and the order matters. Distil first, then prune the student, then fine-tune, then quantize last. Quantizing before pruning means your magnitude rankings are computed on already-rounded weights, which makes them noticeably noisier.
A worked combination on a 440 MB BERT-base: distil to a 6-layer student (66 million parameters, 264 MB, about 97% of the score), structurally prune 20% of the transformer-layer weights with fine-tuning (about 58 million parameters, 232 MB, 96%), then quantize to INT8 (about 58 MB, 95%). Roughly seven to eight times smaller and about four times faster, at a cost of about two points of F1. One caveat on the last step: the 23 million parameters of the embedding table are not a Linear layer, so PyTorch's quantize_dynamic leaves them in FP32 and the file lands nearer 130 MB; ONNX Runtime's quantizer can convert the embedding lookup as well.
The measurement discipline that makes this safe
Every number in this lesson is a typical value, not a guarantee. Your model, your hardware, and your data distribution will produce different ones. So the practical rule is: never adopt an optimisation you have not measured on your own model, your own hardware, and — this is the one people skip — a held-out set drawn from production traffic rather than your original test split.
Measure four things every time, and record them next to the artefact: model size on disk, p50 and p99 latency at your actual batch size, peak memory during inference, and accuracy on the held-out set. Report latency as a percentile, never a mean; quantized models often have a lower median and a comparable tail, and only the tail is what users experience.
Two failure modes are worth naming explicitly because they are quiet. The first is aggregate accuracy hiding a subgroup collapse: a quantized model can lose 0.4% overall while losing 8% on the rarest class, because the rare class relied on activations near the edge of the quantized range. Always break accuracy down per class or per segment, not just in aggregate. The second is benchmarking with a warm cache and a batch size of 64 when production sends batches of one: quantization's benefit comes largely from memory bandwidth, and at batch size 1 with a cold cache the picture can look completely different.
Finally, set the budget before you start. "Under 150 MB, p99 under 200 ms, no class losing more than 2 points" is a specification you can test against and stop at. Without one, optimisation becomes an open-ended activity that ends when someone gets tired, which is exactly when a subgroup regression slips through.