Model Deployment for AI Engineers

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.

Getting BERT-base under the 200 MB limit440 MB92.1100 ms220 MB92.196 ms110 MB90.441 ms110 MB91.741 ms68 MB91.229 msArtefact sizeF1CPU latencyFP32 baselineFP16INT8 post-trainingINT8 QATPruned 40 pct + INT8The 1.7-point drop from post-training INT8 is one outlier activation stretching the scale.
Quantisation-aware training buys back almost all of the accuracy that the naive INT8 cast threw away.

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.

ConstraintTypical requirementWhat violating it looks like
Mobile app bundleUnder ~150 MBStore rejection, or users abandoning the download
Interactive APIp99 under 200 msUsers perceive the product as broken
Edge device RAM512 MB total, sharedOut-of-memory kill mid-inference
Cost per million inferencesSet by unit economicsThe feature is unprofitable and gets cancelled
Battery / thermalSustained load without throttlingDevice 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×10383.4 \times 10^{38}. 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.

TypeBitsRange110M-param modelRelative memory bandwidth
FP3232±3.4e38440 MB1.0×
FP1616±65,504220 MB0.5×
BF1616±3.4e38220 MB0.5×
INT88-128 to 127110 MB0.25×
INT44-8 to 755 MB0.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 ss and a zero-point zz:

q=round ⁣(xs)+zx^=s (q−z)q = \text{round}\!\left(\frac{x}{s}\right) + z \qquad \hat{x} = s\,(q - z)

with

s=xmax⁡−xmin⁡qmax⁡−qmin⁡,z=qmin⁡−round ⁣(xmin⁡s)s = \frac{x_{\max} - x_{\min}}{q_{\max} - q_{\min}}, \qquad z = q_{\min} - \text{round}\!\left(\frac{x_{\min}}{s}\right)

Take a real tensor whose observed range is [-2.5, 3.1], targeting signed INT8 (range -128 to 127).

s=3.1−(−2.5)127−(−128)=5.6255=0.021961s = \frac{3.1 - (-2.5)}{127 - (-128)} = \frac{5.6}{255} = 0.021961

z=−128−round ⁣(−2.50.021961)=−128−(−114)=−14z = -128 - \text{round}\!\left(\frac{-2.5}{0.021961}\right) = -128 - (-114) = -14

Now quantize the value 1.7:

  • q=round(1.7/0.021961)+(−14)=round(77.41)−14=77−14=63q = \text{round}(1.7 / 0.021961) + (-14) = \text{round}(77.41) - 14 = 77 - 14 = 63
  • Dequantize: x^=0.021961×(63−(−14))=0.021961×77=1.6910\hat{x} = 0.021961 \times (63 - (-14)) = 0.021961 \times 77 = 1.6910
  • Error: ∣1.7−1.6910∣=0.0090|1.7 - 1.6910| = 0.0090

The error is bounded by half the scale, s/2=0.01098s/2 = 0.01098, and 0.0090 sits under it. Check the endpoints: x=3.1x = 3.1 gives q=round(141.16)−14=127q = \text{round}(141.16) - 14 = 127, exactly the top of the integer range, and x=−2.5x = -2.5 gives q=−114−14=−128q = -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]:

s=20.0−(−2.5)255=22.5255=0.088235s = \frac{20.0 - (-2.5)}{255} = \frac{22.5}{255} = 0.088235

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=0z = 0 and uses s=max⁡(∣xmin⁡∣,∣xmax⁡∣)/127s = \max(|x_{\min}|, |x_{\max}|)/127. For our original tensor that gives s=3.1/127=0.024409s = 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=0z = 0, a quantized matrix multiply is a plain integer dot product; with a non-zero zero-point, expanding (q1−z1)(q2−z2)(q_1 - z_1)(q_2 - z_2) 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.

Python
import torchfrom torch.ao.quantization import quantize_dynamic# Dynamic: weights quantized ahead of time, activations quantized on the fly.model_int8 = quantize_dynamic(    model, {torch.nn.Linear, torch.nn.LSTM}, dtype=torch.qint8)torch.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:

Python
import torch.ao.quantization as tqmodel.eval()model.qconfig = tq.get_default_qconfig("x86")   # per-channel weights, histogram observermodel_fused = tq.fuse_modules(model, [["conv1", "bn1", "relu1"]])model_prepared = tq.prepare(model_fused)# Calibration: 100-500 REAL samples. No labels needed, no gradients.with torch.no_grad():    for batch, _ in calibration_loader:        model_prepared(batch)model_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\text{round} has zero derivative almost everywhere.

Python
model.train()model.qconfig = tq.get_default_qat_qconfig("x86")model_qat = tq.prepare_qat(tq.fuse_modules(model, [["conv1", "bn1", "relu1"]]))optimiser = torch.optim.SGD(model_qat.parameters(), lr=1e-4)   # 10-100x below originalfor epoch in range(3):                                          # a few epochs, not a full run    for x, y in train_loader:        loss = criterion(model_qat(x), y)        optimiser.zero_grad(); loss.backward(); optimiser.step()model_qat.eval()model_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 PTQStatic PTQQAT
Data neededNone100–500 unlabelled samplesFull labelled training set
Time to applySecondsMinutesHours to days
Typical accuracy drop0.5–2%1–3%0.1–0.5%
Typical CPU speed-up1.5–2×2–4×2–4×
Best forTransformers, RNNsCNNsWhen 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.

Python
from torch.amp import autocast, GradScalerscaler = GradScaler("cuda")for x, y in loader:    optimiser.zero_grad()    with autocast("cuda", dtype=torch.float16):        loss = criterion(model(x), y)      # matmuls in FP16, softmax/norms in FP32    scaler.scale(loss).backward()          # scale up before backward    scaler.step(optimiser)                 # unscale, then step (skips if inf/nan)    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−56.1 \times 10^{-5}. Gradients in a deep network routinely reach 10−810^{-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 = 2162^{16}) before backward(). By the chain rule every gradient is scaled by the same factor, so 1×10−81 \times 10^{-8} becomes 6.55×10−46.55 \times 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).

FP16BF16
Exponent / mantissa bits5 / 108 / 7
Max magnitude65,504~3.4e38
Smallest normal6.1e-5~1.2e-38
Needs loss scalingYesNo
HardwareVolta onward, most GPUsAmpere 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.

Python
import torch.nn.utils.prune as pruneprune.l1_unstructured(model.fc1, name="weight", amount=0.5)   # zero the smallest 50%print((model.fc1.weight == 0).float().mean().item())          # 0.5prune.remove(model.fc1, "weight")   # bake the mask in permanently

Now 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.

Python
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

512×512×3×3=2,359,296 parameters512 \times 512 \times 3 \times 3 = 2{,}359{,}296 \text{ parameters}

Prune 30% of channels throughout the network, so this layer sees 358 inputs and produces 358 outputs:

358×358×9=1,153,476 parameters358 \times 358 \times 9 = 1{,}153{,}476 \text{ parameters}

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.

UnstructuredStructured
What is removedIndividual weightsWhole channels or filters
Accuracy at equal parameter countBetterWorse
Real speed-up on standard hardwareEssentially none below ~90% sparsityDirectly proportional
Disk size after compressionSmaller (zeros compress well)Smaller (fewer values)
Use whenTargeting sparse accelerators or download sizeYou 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.

Python
target, per_round = 0.70, 0.20current = 0.0while current < target:    step = min(per_round, (target - current) / (1 - current))    for module in prunable_layers(model):        prune.l1_unstructured(module, name="weight", amount=step)    fine_tune(model, epochs=2, lr=1e-4)    current = 1 - (1 - current) * (1 - step)    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%1 - 0.8^3 = 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]:

TemperatureClass AClass BClass C
T=1T = 10.99660.002470.00091
T=4T = 40.71590.15970.1244

At T=1T=1 the runner-up carries a gradient of essentially zero. At T=4T=4 — logits become [2.0, 0.5, 0.25], giving e2.0=7.389e^{2.0}=7.389, e0.5=1.649e^{0.5}=1.649, e0.25=1.284e^{0.25}=1.284 over a sum of 10.322 — the fact that B outranks C is a real, learnable signal.

Python
import torch.nn.functional as Fdef distillation_loss(student_logits, teacher_logits, labels, T=4.0, alpha=0.7):    soft = F.kl_div(        F.log_softmax(student_logits / T, dim=1),        F.softmax(teacher_logits / T, dim=1),        reduction="batchmean",    ) * (T * T)                                  # see note below    hard = F.cross_entropy(student_logits, labels)    return alpha * soft + (1 - alpha) * hardteacher.eval()for x, y in loader:    with torch.no_grad():        t_logits = teacher(x)    loss = distillation_loss(student(x), t_logits, y)    optimiser.zero_grad(); loss.backward(); optimiser.step()

The T2T^2 factor is not decoration. Dividing logits by TT scales the gradients of the soft loss by roughly 1/T21/T^2, so without the correction the soft term would contribute 16× less at T=4T=4 than at T=1T=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

TechniqueSize reductionTypical speed-upAccuracy costEffort
Dynamic INT8 PTQ4×1.5–2×0.5–2%Minutes
Static INT8 PTQ4×2–4×1–3%Hours
QAT4×2–4×0.1–0.5%Days
FP16 / BF162×1.5–3× on tensor cores~0%Hours
Structured pruning 30%~1.5–2×~1.4–1.9×1–3%Days
Distillation2–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.