Most writing on distillation is about the training step: which loss, which temperature, how to build the transfer set. In production, what decides whether the project pays off is whether the student you train actually fits the serving budget you have, whether it keeps quality on the traffic you really get, and whether you can ship it without a regression reaching users.
This article treats distillation as a production engineering loop. It starts from the serving budget and derives the student's shape from memory and KV-cache arithmetic, then covers distilling at the precision you will deploy, distilling a draft model for speculative decoding, student-first cascades that defer hard requests to the teacher, re-mining production logs, and the release gates between each stage. The loss functions and the data pipeline have their own pages: distillation from an LLM teacher covers logit caching and a memory-safe KD loss, and the distillation data recipe covers building and verifying the transfer set.
Start from the serving budget, not from the teacher
Do not pick a student from a model zoo and hope. Write down what the deployment must achieve and let that fix the student's size before any training run.
Four numbers are enough: p95 latency targets, the hardware, peak concurrency and the cost ceiling per million requests. These translate into constraints on parameters and on KV cache, because autoregressive decoding at small batch sizes is limited by memory bandwidth, not arithmetic. Each decode step has to stream every weight from memory at least once.
That gives a quick upper bound on single-stream decode speed: tokens per second is at most memory bandwidth divided by the bytes of weights read per token. A model of 8 billion parameters in 16-bit weights is about 16 GB; on an accelerator with roughly 1 TB/s of usable bandwidth the ceiling is about 60 tokens per second per stream before any other overhead. The same arithmetic for a 3B student at 4-bit weights, about 1.5 GB, gives a ceiling several times higher. Real throughput is lower, but the ratio between candidates is a reliable guide.
Choosing the student's shape by arithmetic
Parameter count is not the only knob. Two students with the same parameter count can serve very differently.
- Depth versus width. Layers run sequentially, so a deep narrow model has more serial steps per token and lower utilisation at small batch. For latency-bound serving, prefer fewer, wider layers when quality allows.
- Grouped-query attention. The KV cache scales with the number of key-value heads, not query heads. A student with 8 KV heads instead of 32 holds four times as many concurrent tokens in the same memory.
- Vocabulary. The embedding and output projection can be a large fraction of a small model. If the student must share the teacher's tokenizer, for example to serve as a draft model, that cost is fixed; otherwise a trimmed vocabulary for a narrow domain saves memory and softmax time.
- Context length. Train and serve only the context the product needs.
The KV cache per token is 2 x layers x KV heads x head dimension x bytes per element. Multiply by context length and concurrency to get the memory the cache needs at peak. A short script that does this for each candidate shape, alongside weight memory, makes the decision explicit:
def serving_memory_gb(params_b, weight_bits, layers, kv_heads, head_dim,
kv_bits, context, concurrency):
weights = params_b * 1e9 * weight_bits / 8
kv_per_token = 2 * layers * kv_heads * head_dim * kv_bits / 8
kv = kv_per_token * context * concurrency
return round(weights / 1e9, 2), round(kv / 1e9, 2)
# candidate student: 3B, 4-bit weights, 28 layers, 8 KV heads of 128, fp16 cache
print(serving_memory_gb(3, 4, 28, 8, 128, 16, 4096, 64)) # (1.5, 30.06)The example shows a common surprise: at 64 concurrent 4,096-token requests the KV cache needs twenty times the memory of the 4-bit weights. Shrinking the weights further buys little; reducing context, KV heads or cache precision buys a lot. Decide this before training, because KV heads and layer count cannot be changed afterwards without retraining.
Distil at the precision you will serve
If the student will be served at 4-bit or 8-bit weights, the model that matters is the quantized one, not the 16-bit checkpoint you evaluate in the training notebook. Quantizing afterwards adds a second, unmeasured loss on top of the distillation gap, concentrated on the examples the student was already weakest on.
Two practical options exist. The cheap one is to distil in full precision, quantize with a post-training method, then run a short second distillation phase on the quantized model with the teacher still as the target; this recovers much of the quantization loss for little compute. The thorough one is quantization-aware distillation from the start: the student's forward pass uses fake-quantized weights, gradients flow through a straight-through estimator, and the KD loss compares those outputs with the full-precision teacher. Either way, report numbers from the deployed artefact in the deployed runtime.
# one quantization-aware distillation step (PyTorch-style pseudocode)
with torch.no_grad():
t_logits = teacher(batch.input_ids).logits # full precision teacher
s_logits = student_fakequant(batch.input_ids).logits # weights fake-quantized to int4
T = 2.0
kd = F.kl_div(F.log_softmax(s_logits / T, -1),
F.softmax(t_logits / T, -1),
reduction="batchmean") * T * T
ce = F.cross_entropy(s_logits.flatten(0, 1), batch.labels.flatten(), ignore_index=-100)
loss = 0.7 * kd + 0.3 * ce
loss.backward(); opt.step(); opt.zero_grad()Quantization itself is covered in SLM quantization.
Distilling a draft model for speculative decoding
Sometimes the goal is not to replace the large model but to make it faster. In speculative decoding a small draft model proposes k tokens and the large target model verifies them in one forward pass; the accepted prefix is kept, and the output distribution is provably the target's. The quality of a draft model is measured by one number: the acceptance rate, the probability that the target accepts a drafted token.
If each drafted token is accepted independently with probability a, the expected number of tokens produced per target forward pass is (1 - a^(k+1)) / (1 - a), the formula from the original speculative decoding paper by Leviathan and colleagues. That makes the value of distillation easy to compute:
def tokens_per_verify(a, k):
return (1 - a ** (k + 1)) / (1 - a)
for a in (0.6, 0.7, 0.8):
print(a, [round(tokens_per_verify(a, k), 2) for k in (2, 4, 6)])
# 0.6 [1.96, 2.31, 2.43]
# 0.7 [2.19, 2.77, 3.06]
# 0.8 [2.44, 3.36, 3.95]Raising acceptance from 0.6 to 0.8 at k = 4 lifts tokens per verify step from about 2.3 to about 3.4. Distilling the draft model on the target's own outputs is the direct way to raise acceptance, because it trains the draft to predict what this particular target will say rather than what text in general looks like. The DistillSpec work by Zhou and colleagues studied exactly this and reported that using data generated by the draft model itself (on-policy) and choosing the divergence to suit the decoding strategy both matter. The draft must use the target's tokenizer, and k draft steps must cost a small fraction of one target step. Measure acceptance per traffic slice. The mechanics are in speculative decoding for SLMs.
Student-first cascades with teacher fallback
A student rarely matches the teacher everywhere, but often on most requests. In a cascade every request goes to the student first, and a deferral rule sends the hard ones to the teacher. Cost then depends on the deferral rate, and quality on how well the rule catches the requests the student gets wrong.
The deferral signal can be the student's own confidence (mean token log-probability, or the margin between top choices for classification), a small verifier model, a schema or rule check on the output, or a router that looks only at the input and predicts difficulty before the student runs. Confidence signals need calibration, because students are often overconfident on inputs unlike their transfer set. Calibrate on a held-out set labelled by comparing student and teacher outputs:
def pick_threshold(scores, student_ok, max_defer=0.25):
"""scores: student confidence per example; student_ok: True if output acceptable.
Returns the threshold that minimises bad answers served within a deferral budget."""
pairs = sorted(zip(scores, student_ok)) # lowest confidence first
n = len(pairs)
best = (float("inf"), None)
for i in range(int(n * max_defer) + 1): # defer the i least confident
served_bad = sum(1 for s, ok in pairs[i:] if not ok)
thr = pairs[i][0] if i < n else float("inf")
best = min(best, (served_bad, thr))
return best[1]Two costs are easy to forget. A deferred request pays for both models, so the student must be much cheaper than the teacher for the cascade to win; and a deferred request has higher latency, so the 95th-percentile latency is set by the teacher path.
Re-mining production traffic
Once the student is live, the logs are the real distribution, and the cascade's deferrals list what the student cannot yet do. Each re-distillation round should draw from three pools: a uniform sample of recent prompts so the student tracks drift, the deferred and low-confidence prompts so it learns its weak spots, and prompts tied to negative user signals such as retries, edits or thumbs-down.
On-policy data helps here too. Instead of only training on teacher-written answers, let the student generate, have the teacher score or correct those generations, and train on the corrections; the generalized knowledge distillation work by Agarwal and colleagues frames this as distilling on the student's own sequences. Apply the same privacy controls to mined prompts as to any user data: retention rules, PII scrubbing and per-tenant opt-outs. Deduplicate against the evaluation sets, or the next round's scores will be inflated.
Release gates: offline, shadow, canary
| Gate | What it compares | Blocks release when |
|---|---|---|
| Offline slices | Deployed artefact vs teacher on held-out sets cut by task, language, tenant and length | Any slice drops beyond its tolerance, even if the average improves |
| Shadow | Student runs on mirrored live traffic; outputs logged, never shown | Agreement with the serving model or verifier pass rate falls below target; latency or memory outside budget |
| Canary | A small share of real users get the student or cascade | Guardrail metrics move: retries, escalations, complaint rate, error rate, tail latency |
Slice tolerances matter more than averages. Distillation compresses the long tail first, so a steady average can hide a collapse on a rare language. Building those slices well is covered in SLM evaluation. Keep the router able to send all traffic back to the teacher within minutes.
Worked example: a support-reply model
Take a hypothetical support product where a large hosted model drafts ticket replies, and the budget asks for a 60 percent cost cut with p95 latency under 4 seconds.
The budget arithmetic picks a 3B student with grouped-query attention and 4-bit weights served on the existing GPU pool, with context capped at 4,096. The team mines 200,000 recent tickets, has the teacher write verified replies, and distils with the quantizer in the loop. Offline, the student matches the teacher's rubric score on 88 percent of the held-out set overall but only 71 percent on tickets in Portuguese, so the slice gate fails. The fix is a targeted round of Portuguese prompts rather than a bigger model.
In shadow, a confidence threshold calibrated for a 15 percent deferral budget catches most of the remaining bad drafts. The canary at 5 percent of traffic shows agents editing student drafts no more often than teacher drafts. The blended cost is the student on all traffic plus the teacher on the 15 percent deferred. If the student costs a tenth of the teacher per request, that blend is about 25 percent of the original cost, comfortably inside the 40 percent the budget allows. The deferred tickets become the next round's mining pool.
Failure modes
| Failure | Symptom | Prevention |
|---|---|---|
| Evaluated the wrong artefact | Production quality below offline numbers | Evaluate the quantized model in the serving runtime |
| KV cache ignored | Student fits, but concurrency collapses at peak | Size KV heads, context and cache precision before training |
| Average hides a slice | One tenant or language regresses after rollout | Per-slice tolerances in every gate |
| Overconfident student | Cascade defers too little; bad answers served confidently | Calibrate thresholds on held-out data; recheck after each round |
| Draft tokenizer mismatch | Speculative decoding cannot be used, or acceptance is near zero | Distil the draft with the target's tokenizer |
Trade-offs
| Decision | Option A | Option B |
|---|---|---|
| Replace or accelerate | Student replaces teacher: biggest saving, quality risk | Distilled draft: identical output, smaller saving |
| Cascade or single model | Cascade guards hard requests, adds tail latency | Single student is simpler, needs a higher bar |
| Quantize after or during | After is cheap | During closes the gap at low bit-widths |
What to do next
- Write down the serving budget: p95 latency targets, hardware, peak concurrency and cost per million requests.
- Compute weight and KV-cache memory for each candidate student shape and discard any that cannot meet the budget.
- Pin the teacher version and build the transfer set from mined, privacy-checked production prompts.
- Distil with the deployment quantizer in the loop, and evaluate only the deployed artefact in the serving runtime.
- Define evaluation slices with per-slice tolerances, and run offline, shadow and canary gates in that order.
- If you keep the teacher, decide between a cascade and a distilled draft model, calibrate the deferral threshold or measure acceptance per slice.
- Feed deferrals and negative user signals into the next round, deduplicated against evaluation data.