LLM 蒸馏的显存瓶颈不只在教师模型:离线 Top-K 与分块 KL 把长上下文训练装回单卡
小模型要在低延迟、低成本或本地环境中部署,常见做法不是从零训练,而是先压缩大模型,再用知识蒸馏恢复能力。但蒸馏本身也可能昂贵:在线方案需要让教师与学生同时驻留显存,每一步还要重新运行教师前向计算。新论文《Efficient Knowledge Distillation for LLMs》把问题拆成两个独立瓶颈:教师模型的重复计算,以及语言模型输出头产生的完整词表 logits。
第一刀:把教师移出训练循环
作者先计算教师模型在每个 token 上的 Top-100 概率并缓存,之后学生只读取缓存训练。实验对象是从 Llama 3.1 8B Instruct 得到的约 3.2B 学生模型。在单张 H200、8K 上下文的配置中,离线方案与在线蒸馏给出近似的训练损失曲线;同时峰值显存从约 103GB 降到 78GB,单步时间从 25.9 秒降到 18.5 秒,吞吐从 237 提升到 331 TFLOP/s。论文把收益解释得很直接:教师只运行一次,缓存还能供多组消融实验重复使用。
这里的关键并不是“离线一定优于在线”,而是作者用同一目标函数验证了一个工程交换:把持续占用 GPU 的计算,换成可复用的稀疏教师分布。对于需要反复试验蒸馏配方的团队,这个交换会直接改变实验成本。
第二刀:不要生成完整词表 logits
即使教师离开显存,学生输出头仍会产生随“序列长度 × 词表大小”增长的 logits 张量。作者提出融合分块 KL:按序列块完成输出投影、归一化与稀疏教师项计算,并在反向传播时重算局部结果,因此不需要保存完整 logits。代价是短上下文下会增加一些重算,但显存峰值转为随序列长度线性增长。
在同一张 H200 的真实训练配置中,8K 上下文下,稠密 KL、仅前向分块和融合分块三种实现的峰值显存分别约为 78GB、62GB 和 58GB。到了 32,768 token,稠密方案估算接近 250GB、无法装入单卡;融合分块方案约为 128GB,因此能在 141GB 的 H200 上运行。这也是论文最有现实意义的结果:限制紧凑模型“长上下文恢复训练”的,未必是 Transformer 主体,而可能是最后一层输出与损失计算。
论文还用一个只包含输出投影的受控实验隔离了这个机制。在 32K token 下,融合分块 KL 的峰值显存为 5.45GiB,稠密 KL 为 85.2GiB;在 256K token 下,融合方案使用 11.6GiB,而仅前向分块方案达到 134.2GiB。作者明确提醒,这组数据不是端到端 LLM 训练速度,不能和真实训练结果直接混用。
真正值得拿走的结论
这项工作没有声称发明新的蒸馏算法,而是给出一份可复现的工程配方:**先缓存教师 Top-K 分布,解除教师常驻;再把 KL 损失与输出投影分块融合,解除完整词表 logits 常驻。**消融结果还显示,只有中间层特征损失会让学生性能崩塌,logit 级 KL 是核心;在此基础上再加入隐藏状态特征损失,MMLU 与 GSM8K 才有小幅稳定提升。
边界也很清楚:真实模型实验主要围绕一组 8B 教师与约 3.2B 学生展开,硬件和软件栈集中在 H200、Megatron-Bridge 与 ModelOpt,迁移到其他架构和设备仍需验证。论文与实现链接见原文。
所以,这篇论文最值得关注的不是一个更高分的学生模型,而是一条更实际的提醒:做小模型,不只要压缩参数,也要压缩训练过程里那些本来不必完整存在的中间张量。