Distilling a large teacher into a small student starts with data, and that data is produced by a GPU job. Two shapes of job exist. Sequence-level distillation generates responses from the teacher, which is a decode-heavy sampling workload. Token-level distillation scores fixed text with the teacher and keeps its probability distribution at every position, which is a prefill-only workload with a memory problem at the output layer. The generation job is covered in synthetic data generation on GPUs, and how to choose and filter what the teacher writes is covered in the distillation data recipe.
This page is about the second job: the teacher scoring pass that produces soft targets. It explains why the job is compute-bound, the logits memory wall and the chunked top-k extraction that avoids it, using an inference engine instead, idempotent shards on preemptible GPUs, numerics, the off-by-one that silently ruins datasets, and a fully sized worked example.
A scoring pass is prefill, not generation
A scoring pass runs one forward pass over each prompt-plus-response sequence. There is no autoregressive loop and no KV cache retained across steps. Every token is processed in parallel, exactly like the prefill phase of inference. A forward pass costs about 2N floating-point operations per token for a model with N parameters, so the job is compute-bound at reasonable batch sizes. Its throughput is set by matrix-multiply efficiency, not by memory bandwidth as decoding is.
That changes how you size it. For generation you tune concurrency to fill the KV cache. For scoring you pack sequences into batches with a fixed token budget, keep the tensor cores busy and measure model FLOPs utilisation (MFU): achieved FLOP/s divided by the hardware's dense peak. Long and short sequences mixed carelessly waste compute on padding. Sort by length or pack several sequences into one row with a block-diagonal attention mask.
Architecture of the job
A controller assigns token shards to teacher replicas. Each replica runs the transformer body once per batch and applies the output projection (the lm_head) to slices of positions. For each slice it extracts top-k on the GPU and copies only the compact result to the host. A writer commits each shard atomically, with its manifest written last. The student's data loader later reads token ids and soft targets side by side and verifies they align before computing the distillation loss described in distilling a small model from an LLM.
The logits memory wall and chunked top-k
The output layer is where scoring jobs run out of memory. Logits have shape [positions, vocabulary]. With a 128,256-token vocabulary, one 8,192-token sequence produces 8,192 x 128,256 logits: 2.1 GB in bf16 and 4.2 GB in fp32. A naive log_softmax materialises another full copy, and a batch of eight sequences then needs tens of gigabytes for one layer's output, on top of the teacher's weights.
The fix is to never hold all positions' logits at once. Run the transformer body to get hidden states, which are small (positions x model width), then apply lm_head to a slice of positions, reduce that slice to top-k, discard its logits and move on. Normalisation does not need the full log_softmax: log pi = zi - logsumexp(z), so take top-k on raw logits and subtract the slice's logsumexp, computed in fp32.
import torch
K = 32
CHUNK = 1024 # positions per lm_head slice: 1024 x 128,256 x 4 B = 525 MB in fp32
@torch.no_grad()
def score_sequence(model, input_ids, attention_mask, loss_mask):
"""Teacher-forced top-k log-probs at the positions the student trains on.
All inputs are [1, T]. The distribution at position t predicts token t+1,
so the loss mask and the targets are shifted by one.
"""
hidden = model.model(input_ids=input_ids,
attention_mask=attention_mask,
use_cache=False).last_hidden_state # no KV cache: one pass
keep = loss_mask[0, 1:].bool() # target t+1 is a response token
hidden = hidden[0, :-1][keep] # [N, d]
targets = input_ids[0, 1:][keep] # stored for alignment checks
ids, logp, tail = [], [], []
for start in range(0, hidden.shape[0], CHUNK):
z = model.lm_head(hidden[start:start + CHUNK]).float() # [c, V]
lse = torch.logsumexp(z, dim=-1, keepdim=True)
top_z, top_ids = z.topk(K, dim=-1)
top_logp = top_z - lse # normalised without a full log_softmax
ids.append(top_ids.to(torch.int32).cpu())
logp.append(top_logp.to(torch.float16).cpu())
tail.append((1 - top_logp.exp().sum(-1)).clamp_min(0).to(torch.float16).cpu())
return targets.to(torch.int32).cpu(), torch.cat(ids), torch.cat(logp), torch.cat(tail)This follows the Hugging Face layout, where model.model is the transformer body and model.lm_head the output projection. Check your architecture's forward method before copying it. Some models transform logits after the projection, for example final-logit soft-capping in Gemma 2, and calling lm_head directly skips that step. Copying top-k rather than full logits to the host also cuts PCIe traffic from about 256 KB per token in bf16 to under 200 bytes.
Scoring with an inference engine
You can also score with an inference engine. In vLLM, pass the prompt and response together as the prompt, set max_tokens=1, and request prompt_logprobs to get the top entries at every prompt position. Two limits apply. The engine caps the number of log-probs per position with max_logprobs, which defaults to 20, so k = 32 requires raising it when constructing LLM. Results also include the actual token's log-prob, so up to k + 1 entries can come back.
from vllm import LLM, SamplingParams
llm = LLM(model=TEACHER, tensor_parallel_size=4, max_logprobs=32)
params = SamplingParams(max_tokens=1, prompt_logprobs=32)
for out in llm.generate([{"prompt_token_ids": ids} for ids in batch], params):
rows = out.prompt_logprobs # one entry per prompt position (the first is None),
# each a dict of token id -> LogprobThe engine gives you batching, tensor parallelism and memory management for free. It returns Python dictionaries per position, so converting billions of positions into arrays can make the host CPU the bottleneck. Measure that in your pilot before committing. The budget below assumes a tensor-parallel replica, from the engine or a TP-capable loader: splitting layers across GPUs that run one after another leaves most of them idle. If you also sample with the engine, check in your version's documentation whether reported log-probs are taken before or after temperature and other logit processing. For soft targets you want the raw distribution.
Idempotent shards on preemptible GPUs
Scoring jobs are ideal for preemptible or spot GPUs: there is no cross-shard state, and losing a shard costs only that shard's compute. That works only if writes are idempotent, meaning a shard that was half-written when its machine vanished must never be mistaken for a finished one.
import hashlib, json, os, shutil
import numpy as np
def shard_done(out_dir, shard_id):
return os.path.exists(os.path.join(out_dir, shard_id + ".manifest.json"))
def commit_shard(out_dir, shard_id, arrays, meta):
tmp, final = os.path.join(out_dir, "." + shard_id + ".tmp"), os.path.join(out_dir, shard_id)
shutil.rmtree(tmp, ignore_errors=True)
os.makedirs(tmp)
sha = {}
for name, arr in arrays.items(): # targets, ids, logp, tail
path = os.path.join(tmp, name + ".npy")
np.save(path, arr)
with open(path, "rb") as f:
sha[name] = hashlib.sha256(f.read()).hexdigest()
shutil.rmtree(final, ignore_errors=True) # leftovers of a preempted attempt
os.replace(tmp, final)
manifest = dict(meta, shard=shard_id, rows=int(len(arrays["targets"])), sha256=sha)
mtmp = os.path.join(out_dir, shard_id + ".manifest.tmp")
with open(mtmp, "w") as f:
json.dump(manifest, f)
os.replace(mtmp, os.path.join(out_dir, shard_id + ".manifest.json")) # commit markerThe manifest is the commit record. It is written last, and the controller skips any shard that has one. Put in meta everything that defines the targets: teacher checkpoint hash, tokenizer hash, chat template version, k, dtypes, CHUNK, the loss-mask rule and the code commit. On object stores without atomic rename, upload the arrays first and the manifest object last. The same rule then holds.
Numerics
- Compute logsumexp in fp32. In bf16 its rounding error shifts every log-prob in a row by the same amount, which corrupts the targets without raising an error.
- Store log-probs, not probabilities. fp16 probabilities underflow for rare tokens. fp16 log-probs keep about three significant digits over the useful range. Clamp any -inf from masked vocabulary entries to a finite floor before casting.
- Do not expect bitwise reproducibility. Changing the tensor-parallel degree, batch composition or kernel changes reduction order, so log-probs differ in low bits between runs. Test with tolerances, for example top-1 agreement above 99.9% and mean absolute log-prob difference below 0.01, not with checksums of the arrays.
- Keep the tail mass. It tells the student loss how much probability the top-k dropped, and rows with large tails mark exactly the uncertain positions worth learning from.
Worked example: soft targets for a billion tokens
Suppose you need soft targets for a corpus of 1 billion tokens of prompt plus response, of which 600 million are response positions the student trains on. The teacher has 70 billion parameters and a 128,256-token vocabulary. The hardware is 80 GB H100 SXM GPUs, eight per node.
| Quantity | Arithmetic | Result |
|---|---|---|
| Teacher weights (bf16) | 70e9 x 2 B | 140 GB, so at least 2 GPUs |
| Replica layout | TP = 4: 320 GB total, 180 GB headroom | 2 replicas per node |
| Compute | 2 x 70e9 x 1e9 tokens | 1.4e20 FLOP |
| Effective throughput per GPU | 989 TFLOP/s dense bf16 peak x 40% MFU (assumed) | about 396 TFLOP/s |
| GPU time | 1.4e20 / 3.96e14 | about 98 GPU-hours |
| Wall clock on 8 nodes | 98 / 64 GPUs | about 1.5 hours plus overheads |
| Stored per position | 32 x 4 + 32 x 2 + 2 + 4 bytes | 198 bytes |
| Total soft targets | 600e6 x 198 B | about 119 GB |
| Full bf16 logits instead | 600e6 x 128,256 x 2 B | about 154 TB |
The 40% MFU figure is an assumption to replace with your pilot's measurement. Run 1% of the shards, compute achieved FLOP/s from tokens per second, and re-derive the budget. The headroom in a TP-4 replica goes to activations for large packed batches. Since scoring keeps no KV cache, it is usually better spent on batch size than on more replicas. The last two rows are why on-GPU top-k is not optional: 119 GB fits on one disk, while 154 TB is a storage project. At student training speed, reading 198 bytes per token is a trivial load for any object store.
Failure modes
- Off-by-one alignment. Logits at position t predict token t+1. Storing them against token t trains the student to predict the current token, and the loss still goes down. Store the targets and assert in the loader that the student's shifted labels equal them, shard by shard.
- Scoring prompt tokens by accident. If the loss mask is not shifted with the targets, the first response token is dropped and the last prompt token is kept.
- Template drift. Scoring with one chat template and training with another misaligns every sequence. Hash the rendered template into the manifest and check it at load time.
- Silent truncation. Sequences longer than the teacher's maximum context are cut or rejected depending on the engine. Count tokens per shard against the source and fail on mismatch.
- Half-written shards after preemption. Without a manifest-last commit, a resumed job either redoes everything or trusts a truncated array.
- Stale teacher. Rerunning part of a corpus after a checkpoint update mixes two teachers. The manifest's checkpoint hash catches it if the loader checks it.
- Skipped logit transforms. Applying
lm_headdirectly on a model that soft-caps or scales logits in its forward method yields a different distribution from the real teacher.
Trade-offs
Offline cache or online teacher. Caching pays the teacher forward pass once and makes every later student experiment cheap. Online scoring avoids storage but pays again every epoch and every rerun, and it is required for on-policy distillation, where the teacher scores the student's own samples. Hand-rolled or engine. A hand-rolled PyTorch scorer gives control over chunking, dtypes and output format. An engine gives throughput tooling but hands results back as Python objects. k. Larger k keeps more of the distribution at linear storage cost. Choose it from the measured tail-mass distribution. Replicas versus TP degree. Use the smallest TP degree that fits weights plus a large batch, because more replicas scale scoring almost linearly while larger TP groups add communication. After scoring, the dataset feeds a normal fine-tuning job, with the KD loss alongside or instead of cross-entropy.
What to do next
- Decide token-level or sequence-level distillation. If sequence-level, use the generation job and skip most of this page.
- Implement the chunked scorer, and assert on a few sequences that its top-1 ids and log-probs match a full
log_softmaxwithin tolerance. - Write the targets array and an alignment assertion in the student loader before scoring anything at scale.
- Run a 1% pilot, measure tokens per second and MFU, and redo the budget table with your numbers.
- Measure tail mass on the pilot and choose k from it.
- Use manifest-last commits, put every target-defining setting in the manifest, and kill a worker mid-shard to prove resumption works.
- Schedule the full run on preemptible capacity, then verify shard count, row counts and hashes before the student trains.