Offline Hidden-State Distillation from Compact Teacher Caches
Abstract
For language models, logit distillation is now routinely performed offline: the teacher's next-token distributions are computed once, ahead of training, and then cached for use in subsequent training runs. Storing the full distributions would be prohibitive, but truncation or sampling compresses the cache by orders of magnitude without significantly impacting performance. Hidden-state distillation, which additionally uses the teacher's hidden states to supervise the student's intermediate computation, has no compact offline counterpart. Existing methods either perform teacher inference within the training loop, slowing down training significantly, or require enormous amounts of storage. We show that a per-token random projection of the teacher's hidden states, coarsely quantized, lets hidden-state distillation run from a cache of the same order of size as the logit cache, at the training throughput of logit distillation. With the hidden states compressed 32, students distilled from four teachers of 1.5B to 8B parameters need at most about 2% more training tokens than students trained on the full hidden states to reach the same held-out loss at 1B tokens. Over training runs of up to 15B tokens, they need about 1% more tokens than the full-state students and match their downstream accuracy, while a student distilled from logits alone needs 1.4 to 1.7 times as many.
est. 32% chance this paper gets accepted at ICLR 2027.
What do you think this paper will get?
All positions stay anonymous.