离线缓存教师 logits 做蒸馏,每轮快 29%,单卡上下文翻四倍到 32K

Efficient Knowledge Distillation for LLMs: Offline Top-K Logits and a Fused Chunked KL Loss

Bakbergen Ryskulov, Iker García-Ferrero, David Montero, David Jansen, Ali Hashemi, Jezabel R. Garcia, Antonio Tiene, Román Orús

cs.CL, cs.AI, cs.LG

2026-08-04

用 Llama-3.1-8B 蒸馏出约 3.2B 学生:把教师 top-100 logits 离线缓存、训练时不再加载教师,每轮迭代快 29%、吞吐高约 40%;再配一个不实例化全词表 logit 张量的分块 KL 损失,单卡 H200 能训到 32K 上下文,学生 MMLU 60.6%、GSM8K 67.5%。

这篇在解决什么

小模型往往是延迟、成本、私有部署约束下唯一能用的选择,但它通常不是从头训出来的:知识蒸馏(KD)把一个大模型的能力压进小模型,这步压缩基本决定最终效果。然而它很贵。常规做法里教师和学生要同时在显存里,算 KL 散度还得先把 logit 升到整个词表大小(动辄十几万维),这个全词表张量按序列长度放大,显存一爆就把能训的上下文长度卡死了。

这是一篇面向实践的工程研究,围绕两个系统层面的改动把蒸馏做得又快又省显存。作者来自做量子化压缩 Llama 的 CompactifAI 团队,代码开源。

方法

两个组件,各自独立可用。

第一,离线 top-K 蒸馏。常规在线蒸馏每个 batch 都让教师前向一遍、算出它的全分布当目标。这里改成一次性把教师每个位置概率最大的 top-100 logits 缓存下来,训练学生时只对着这份稀疏缓存算 KL,教师从此不进训练循环、不占显存。top-100 而非全词表,是工程上省事的关键。

第二,fused 分块 KL 损失。标准 KL 要先把学生隐藏状态过输出头、实例化出「序列长度 × batch × 词表大小」的巨大 logit 张量再算损失。这里把输出投影直接融进损失核、按序列分块处理,每块算完立刻丢弃,峰值显存从随序列长度平方增长降到线性。那个卡死上下文长度的显存尖峰就没了。

结果

在线与离线对比(8K 上下文,单卡 H200):训练损失几乎一致;显存从 103GB 降到 78GB;每轮迭代从 25.9 秒降到 18.5 秒(快约 29%);吞吐从 237 升到 331 TFLOP/s(约 40%)。

加上分块 KL 损失:显存再从 78GB(密集)降到 58GB(fused);单卡能训到 32K 上下文,约是密集做法的四倍。

隔离损失核的 toy 基准(4K–256K):32K 下分块损失只用 5.45GB,密集要 85.2GB,省 15.6 倍;256K 下速度是前向分块的 3.3 倍(0.630 对 0.190 迭代/秒)。

模型质量(教师 Llama-3.1-8B,学生约 3.2B):

损失配置MMLUGSM8K
仅中间层特征损失约 28%约 4%
仅 logit KL59.9%65.9%
logit KL + 隐藏态特征损失60.6%67.5%

学生在不到教师一半参数量下保住了大部分短上下文能力,与教师在 MMLU 上差约 9 个点(WinoGrande、GSM8K 差距更大)。

为什么重要

对做模型压缩、私有化部署的人,这是一份可直接抄的配方:离线缓存把教师踢出训练循环,长上下文蒸馏不再需要堆卡。把上百次消融跑得起,正是这种高效带来的附带好处。结论也简单:logit KL 必须有,再叠一个隐藏态特征损失能稳定涨一点。

局限与存疑

作者自己列了几条:只验证了一个师生对(8B 到 3.2B),换架构是否成立未测;256K 的数据来自隔离的 toy 网络而非端到端训练;结果在 Megatron-Bridge 加 ModelOpt 加 H200 上得到,别的框架没验证。还有一点:学生与教师在 MMLU、GSM8K 上仍差 9 分以上,说明这套高效配方主要解决「训得动」,不是「追平教师」,压缩比的极限这篇文章没有回答。

术语

原文与代码

社区讨论

相关论文

全部论文解读