Making Knowledge Distillation Cheap Enough to Run at Scale

2026-08-20 · Hugging Face

Making Knowledge Distillation Cheap Enough to Run at Scale

Knowledge distillation — training a smaller student model to match the performance of a larger teacher — is a well-established technique in machine learning. With the recent wave of open-source LLMs such as gpt-oss, Qwen, GLM, and Kimi, it has once again become a mainstream research topic.

The Cost of Deploying Giant Models

Deploying these models is extremely expensive. The Kimi-K3 model, for example, has 2.8 trillion parameters and requires roughly 3TB of VRAM just to load. Consequently, compressing large models and recovering capabilities via knowledge distillation has become standard practice. Companies like Nvidia (Nemotron 3 Puzzle 75B) and Multiverse Computing (Hypernova 60B) have recently released high-quality distilled models.

Why Distillation Is So Expensive

While the distillation step largely determines final model quality, it is also usually the most expensive part of the pipeline. The standard online setup using KL divergence keeps both teacher and student in memory simultaneously. At every training step the teacher must run a full forward pass to produce a probability distribution over the entire vocabulary, which the student then matches.

This is memory-intensive. For gpt-oss-120b (vocabulary size 201,088), a sequence length of 32K and batch size of 4 produces a teacher-probability tensor of shape 4 × 201,088 × 32,768. In bfloat16 that single tensor already consumes about 50GB. Adding gradients, activations, weights, and optimizer states pushes peak VRAM to roughly 250GB — beyond the capacity of even an H200 GPU.

Two Systems Innovations

The paper *Efficient Knowledge Distillation for LLMs: Offline Top-K Logits and a Fused Chunked KL Loss* tackles this with two key changes.

Offline Distillation

Instead of recomputing the teacher on every step, the method runs the teacher once, caches the top-100 most likely logits per position, and trains the student against that cache. The teacher never needs to sit in memory during training, and the cache can be reused across many experiments.

A Memory-Efficient KL Loss

The authors compare three mathematically equivalent ways to compute the KL loss:

  • Dense KL: Reconstructs a full dense teacher probability grid from the cached top-100 logits and compares it with the student’s dense log-probabilities. This is closest to standard online distillation but requires holding the entire vocabulary × sequence matrix in memory twice.
  • Forward-chunked KL: Keeps the teacher sparse (only cached top-100 logits) and computes the loss in sequence chunks. It removes the dense teacher matrix and is the fastest method, but the student’s full logits grid must still be materialized for the backward pass.
  • Fused Chunked KL (main contribution): Fuses the model’s output projection directly into the loss computation. It never materializes the student’s full logits grid. Instead, it processes one chunk of the sequence end-to-end, projects hidden states to logits for that chunk, folds the result into the running loss, and discards the chunk. The backward pass recomputes each chunk on the fly. Although the projection is performed twice (forward and backward), peak memory usage drops dramatically.

Impact

Combined, these changes cut VRAM requirements from ~250GB spikes to ~128GB peaks. This makes long-context healing possible on a single GPU and renders large-scale distillation experiments practical. The techniques represent a significant reduction in cost compared with default implementations in PyTorch or Megatron.

Source