Distillation trains a small student model to imitate a larger teacher. The loss functions are well known and are covered, with token-level versus sequence-level objectives, on-policy variants and evaluation, in SLM distillation architecture. This article is about everything around the loss when the teacher is a large LLM: how much of the teacher you can see, whether you run it during training or once in advance, how to store its predictions without filling a data centre, how to compute the loss without running out of accelerator memory, and what to do when teacher and student do not share a tokenizer.

These are the decisions that set the budget and the ceiling of a distillation project. Get them wrong and you either pay for the teacher many times over or throw away most of the signal that made distillation worth doing.

Advertisement

Why an LLM teacher changes the engineering

With classic distillation the teacher is perhaps a few times larger than the student and both fit on one device. An LLM teacher can be fifty or a hundred times larger. Its forward pass then dominates the compute bill, its full output distribution per token is a vector as long as the vocabulary, often over 100,000 entries, and it may be reachable only through someone else's API. Each of these facts pushes the design in a particular direction.

The signal is worth the trouble. A one-hot next-token label says only which token came next. The teacher's distribution also says which alternatives were plausible and how plausible, which is far more information per token. Several open small-model families have reported training with teacher distributions from larger siblings for this reason, although details of their pipelines vary and are only partly published.

Three levels of teacher access

AccessWhat you get per tokenPossible objectivesTypical situation
White-box: weightsFull logits, hidden states if wanted.Full-vocabulary KL, top-k KL, feature matching, on-policy distillation.Open-weight teacher you can run.
Grey-box: top-k log-probsA few highest-probability tokens with log-probabilities, on generated text.Sparse KL on the returned tokens plus sequence-level training.Hosted API that exposes log-probabilities.
Black-box: text onlySampled responses.Sequence-level distillation: fine-tune on teacher outputs.Hosted API without log-probabilities.

Sequence-level distillation from text works well for teaching formats, tasks and styles, and it is the only option for a black-box teacher. Token-level distillation transfers more of the teacher's uncertainty and usually does better per training token, but needs at least grey-box access and, in practice, a shared tokenizer. If you have white-box access, you choose between them; if you do not, the access level has chosen for you.

Advertisement

The pipeline

Prompt corpustask prompts + seed textTeacher inference farmlarge LLM, bf16, batchedLogit cachetop-k ids + log-probsStudent trainerchunked KD + CE lossCheckpointsstudent weightsTeacher textresponses (black-box)Evaluationheld-out tasks, teacher gappromptsshardssparse targetssampled textsequence-levelsaveloadfix data gapsoffline teacher: pay the teacher forward pass once, reuse the cache every epoch
A distillation pipeline with an offline teacher: prompts go through the teacher once; the student trains against cached sparse targets and sampled text; evaluation feeds back into the prompt corpus.

The prompt corpus decides what the student learns. For general capability it is pretraining-style text; for a task specialist it is task prompts, ideally drawn from real traffic and deduplicated against your evaluation set. The teacher farm runs batched inference and writes shards. The trainer streams shards next to the token ids. Evaluation measures the gap to the teacher on held-out tasks and tells you where the prompt corpus is thin.

Online or offline teacher: do the arithmetic

A forward pass costs about 2N floating-point operations per token for a model with N parameters, and a training step, forward plus backward, about 6N. Take a 70-billion-parameter teacher and a 1-billion-parameter student. The teacher's forward pass costs about 140 GFLOP per token and the student's training step about 6 GFLOP per token. Running the teacher online costs over twenty times more compute than training the student, and you pay it again every epoch and every time you rerun an experiment with a different learning rate.

Memory tells the same story. In bf16 the teacher's weights alone take about 140 GB, so an online teacher needs its own set of accelerators, tensor-parallel across several devices, sitting next to the student's training job and keeping pace with it. An offline teacher pays the forward pass once, in a plain inference job that can use cheaper, preemptible capacity, and makes every later experiment cheap.

Online wins in two cases. On-policy distillation, where the student generates text and the teacher scores it, needs the teacher at training time by definition. And if the prompt corpus is so large that you pass over it only once, caching saves no compute and costs storage. Everything else, including most task-specialist work, should cache.

Caching top-k logits

Storing full distributions is out of the question. A vocabulary of 128,256 entries in fp16 is about 256 KB per token, so a billion tokens would need about 256 TB. Storing only the top k entries changes the picture. With k = 32, each token needs 32 int32 ids (128 bytes) and 32 fp16 log-probabilities (64 bytes), 192 bytes in total, so a billion tokens take about 192 GB. Use int32 for ids: modern vocabularies exceed the 65,535 that a uint16 can hold, and silently wrapped ids corrupt the targets without any error.

Top-k keeps most of the probability mass because LLM distributions are sharply peaked at most positions. Record the tail mass, one minus the sum of the kept probabilities, so you can see where truncation loses information; positions with large tails are exactly the uncertain ones you most want to learn from, so pick k by measuring the tail distribution on a sample, not by habit.

# Offline teacher pass: store top-k log-probs per token as memory-mapped shards.
import numpy as np, torch

K = 32
@torch.no_grad()
def cache_shard(teacher, batches, path, n_tokens):
    ids_mm  = np.lib.format.open_memmap(path + ".ids.npy",  "w+", np.int32,   (n_tokens, K))
    logp_mm = np.lib.format.open_memmap(path + ".logp.npy", "w+", np.float16, (n_tokens, K))
    tail_mm = np.lib.format.open_memmap(path + ".tail.npy", "w+", np.float16, (n_tokens,))
    pos = 0
    for input_ids, attn in batches:                       # packed sequences, same tokenizer as student
        logits = teacher(input_ids=input_ids, attention_mask=attn).logits.float()
        logp = torch.log_softmax(logits, dim=-1)          # [B, T, V]  (process in slices if V is large)
        top_logp, top_ids = logp.topk(K, dim=-1)
        keep = attn.bool()
        top_ids, top_logp = top_ids[keep], top_logp[keep] # drop padding: [N, K]
        tail = 1.0 - top_logp.exp().sum(-1)               # probability mass outside the top-k
        n = top_ids.shape[0]
        ids_mm[pos:pos + n]  = top_ids.cpu().numpy()      # int32: vocab ids exceed 65,535
        logp_mm[pos:pos + n] = top_logp.cpu().numpy()
        tail_mm[pos:pos + n] = tail.clamp_min(0).cpu().numpy()
        pos += n
    return pos

Shard the cache in step with the token shards, write a manifest with the teacher checkpoint, tokenizer hash, k and the sampling settings, and verify on load that the student's tokenizer hash matches. A cache built with the wrong tokenizer produces a loss that decreases and a model that is worse.

A memory-safe KD loss

The naive loss materialises the student's logits for the whole batch. With 8 sequences of 4,096 tokens and a 128,256-entry vocabulary, that is 4.2 billion values, about 16.8 GB in fp32, before autograd keeps a second copy for the backward pass. On an accelerator that is also holding weights, optimizer state and activations, this single tensor is often what causes out-of-memory errors.

The fix is to compute the loss in chunks of positions straight from the final hidden states, and to recompute each chunk's logits during the backward pass instead of storing them. Peak memory for the loss becomes one chunk's logits, for example 2,048 positions times the vocabulary, a few hundred megabytes. The extra compute is one more output-projection matrix multiply per chunk, which is small next to the rest of the model.

import torch, torch.nn.functional as F
from torch.utils.checkpoint import checkpoint

def _chunk_loss(h, W, top_ids, top_logp, labels, T, alpha):
    logits = (h @ W.t()).float()                          # [c, V] only for this chunk
    logq = F.log_softmax(logits / T, dim=-1)
    q_top = logq.gather(-1, top_ids)                      # student log-probs at teacher's top-k
    p_top = F.softmax(top_logp.float() / T, dim=-1)       # renormalise teacher over its top-k
    kd = (p_top * (p_top.clamp_min(1e-9).log() - q_top)).sum(-1)       # KL(p || q) on top-k
    ce = F.cross_entropy(logits, labels, reduction="none")              # ground-truth next token
    return (alpha * T * T * kd + (1 - alpha) * ce).sum()

def distill_loss(hidden, lm_head_weight, top_ids, top_logp, labels, T=1.0, alpha=0.5, chunk=2048):
    """hidden: [N, d] final hidden states of real (non-pad) positions; never builds [N, V]."""
    total = hidden.new_zeros((), dtype=torch.float32)
    for s in range(0, hidden.shape[0], chunk):
        e = s + chunk
        total = total + checkpoint(_chunk_loss, hidden[s:e], lm_head_weight,
                                   top_ids[s:e], top_logp[s:e], labels[s:e], T, alpha,
                                   use_reentrant=False)   # recompute chunk logits in backward
    return total / hidden.shape[0]

Three details matter. The T squared factor keeps gradient magnitudes comparable across temperatures, as in the original formulation by Hinton and colleagues. Applying temperature to cached top-k log-probabilities and renormalising over k is an approximation of tempering the full distribution; it is usually fine near T = 1 and degrades at high temperatures, so cache at the temperature you intend to train with. And mixing in ordinary cross-entropy on the true next token, the (1 - alpha) term, anchors the student where the teacher is wrong or the top-k set misses the label.

When the tokenizers differ

Token-level KL needs the teacher and student to predict over the same vocabulary at the same positions. The simplest solution is to give the student the teacher's tokenizer, and if the vocabulary is too large for a small model, to trim it carefully afterwards, as discussed in tokenizer vocabulary trimming. Deciding this before pretraining saves a great deal of trouble.

If the tokenizers must differ, you have three options. Sequence-level distillation needs no alignment at all: decode the teacher's text and retokenize it for the student. Span alignment maps both tokenizations back to character offsets and applies the KL only where token boundaries coincide and the tokens are identical strings, which recovers signal on common words and loses it on the rest. Research methods that compare distributions without matching vocabularies exist, but they are younger and less predictable; measure them against the sequence-level baseline before relying on them.

A worked budget

Suppose you are distilling a 70B teacher into a 1B student for customer-support tasks on 2 billion tokens of prompts and teacher responses, with a shared tokenizer. The offline teacher pass is 2 billion tokens times 140 GFLOP, about 2.8 x 10^20 FLOP. At an assumed sustained 400 TFLOP/s per accelerator, that is about 7 x 10^5 accelerator-seconds, roughly 200 accelerator-hours. The cache at k = 32 is about 384 GB. Training the student for three epochs is 3 x 2 billion x 6 GFLOP, about 3.6 x 10^19 FLOP, under 30 accelerator-hours at the same rate. The same three epochs with an online teacher would add about 600 accelerator-hours.

Treat throughput figures as assumptions to replace with your own measurements; real utilisation depends on sequence length, batch size and kernels. The training loop around the loss, including the optimizer, schedule and checkpointing, is the same as in the SLM pretraining loop.

Failure modes

SymptomCauseFix
Loss falls, quality dropsTokenizer or vocabulary mismatch between cache and student.Store and check the tokenizer hash; assert id ranges on load.
Out of memory in the lossFull [B, T, V] logits materialised.Chunked loss with recomputation.
Student copies teacher mistakesPure KD with no ground-truth anchor.Mix in cross-entropy; filter teacher outputs that fail verification.
Good on training prompts, poor in productionPrompt corpus unlike real traffic.Sample prompts from logs; deduplicate against evaluation.
Repetitive or degenerate generationExposure bias from purely offline targets.Add an on-policy phase; see the SLM distillation article.
Silent garbage targetsIds stored as uint16.Use int32 ids.

Licences and governance

Check the teacher's licence or terms of service before you start. Some hosted services and some open-weight licences restrict using outputs to train other models, particularly competing ones, and the restrictions differ between providers and versions. Record the teacher, its version and its terms in the cache manifest so the provenance of every student checkpoint is traceable. Treat the prompt corpus as training data with the same privacy review as any other, because the student can memorise it.

Finally, judge the student on the tasks it will serve, with the methods in SLM evaluation, and report the gap to the teacher rather than only absolute scores.

What to do next

  1. Decide your access level and confirm the licence allows distillation.
  2. Make the student share the teacher's tokenizer if you still can.
  3. Do the FLOP arithmetic for online versus offline; cache unless you need on-policy data or a single pass.
  4. Measure the top-k tail mass on a sample and choose k from it; store ids as int32 with a manifest.
  5. Implement the chunked, recomputed KD loss and confirm peak memory with a profiler.
  6. Mix cross-entropy with KD and tune alpha and temperature on a small run.
  7. Evaluate on production-like tasks and report the gap to the teacher.
Key takeaway: Distilling from a large LLM is mostly systems work around a simple loss. Your access level decides which objectives are possible; the FLOP arithmetic almost always favours running the teacher once and caching top-k log-probabilities; a chunked, recomputed loss keeps full-vocabulary logits out of memory; and a shared tokenizer is the cheapest decision you can make early. Keep a ground-truth term in the loss, record provenance and licences in the cache manifest, and measure the student on the work it will actually do.