HF博客介绍如何降低知识蒸馏的计算成本,使蒸馏模型能在生产环境大规模部署,包括具体工程思路和实验数据。
知识蒸馏(knowledge distillation),即训练一个小的 student 模型去匹配大的 teacher 模型的表现,是机器学习中一项广为人知的技术。随着最近开源大语言模型(如 gpt-oss、Qwen、GLM 或 Kimi)的浪潮,这一技术再次成为主流研究方向。部署这些超大模型代价高昂:最新的 Kimi-K3 模型拥有 2.8 万亿参数,仅加载就需要大约 3TB 的 VRAM。因此,将它们压缩成更小的模型,并通过知识蒸馏恢复原有能力已成为标准做法,英伟达(Nemotron 3 Puzzle 75B)或 Multiverse Computing(Hypernova 60B)等公司最近都发布了高质量的压缩模型。
蒸馏步骤决定了最终质量的大部分,但它通常也是整个流程中最昂贵的部分。同时在内存中保持 teacher 和 student 两个模型,并为每个 token 在整个词表上产生概率分布,需要消耗大量的 VRAM,通常只有在数百个 GPU 上并配合精细的 tensor 并行策略才能实现。我们最新的论文 Efficient Knowledge Distillation for LLMs: Offline Top-K Logits and a Fused Chunked KL Loss,通过两个系统层面的改进来解决这一问题:缓存 teacher 的 top-K logits 一次,这样 teacher 就无需与 student 同时驻留内存;以及一种新的、内存高效的 KL 散度损失函数,避免实例化整个词表大小 × 序列长度的矩阵,将 VRAM 消耗降低到 PyTorch 或 NVIDIA Megatron-Bridge 等库的默认实现所能达到的水平以下。这两个改进加在一起,大幅降低了训练成本,使得在单 GPU 上进行长上下文修复(long-context healing)成为可能,也让大规模实验变得切实可行。
标准做法是使用 Kullback-Leibler 散度损失(KL loss)进行在线蒸馏,同时在内存中保持 teacher 和 student 两个模型。在每个训练步中,teacher 运行完整的 forward pass 来产生输出分布,然后训练 student 去匹配它。这是表达力最强的设置,因为完整的 teacher 分布是可用的,但同时也是内存和计算最密集的:每个 token 位置必须持有两个完整词表的 tensor,而且 teacher 必须在每一步都重新计算,尽管它的行为在整个训练过程中不会改变。
作为一个实际例子,gpt-oss-120b 的词表有 201,088 个 token。在序列长度 32K、batch size 4 的情况下,仅 teacher-probability tensor 的形状就是 4 × 201,088 × 32,768;以 bfloat16 存储,仅这一个 tensor 就需要约 50GB 的 VRAM。加上梯度、激活值、模型权重和优化器状态,单次蒸馏训练迭代的峰值 VRAM 约为 250GB,甚至超过了一张 H200 或 B200 GPU 的容量。在这篇文章中,我们展示了重新 formulation KL loss 以分块处理数据,可以将这一成本降低到几乎为零。
Dense KL 峰值约为 250GB,超出单张 H200 的 141GB 容量。fused chunked loss 永远不会形成那个峰值,峰值约为 128GB。来源:论文图 1。

离线蒸馏。 与其在每一步都重新计算 teacher,我们只计算一次输出,缓存每个位置最可能的 100 个 token,然后让 student 针对这个缓存进行训练。在训练期间 teacher 永远不需要驻留内存,一旦缓存存在就不需要再次运行,因此同一个缓存可以跨多次消融实验复用。
Fused, chunked KL loss。 要理解为什么 loss 本身也很昂贵,先设想它实际构建的是什么:对于序列中的每个 token 位置和词表中的每个词,都需要一个数字来描述 student 的预测与 teacher 的分歧程度。如果把它摊开成一张网格,就是词表每项一行、序列位置每列一组;对于一个包含 100K+ 词且序列很长的词表,这张网格是巨大的,而默认的 KL loss 计算方式在产生任何一个数字之前就必须构建完整的网格。
我们比较了三种计算同一 loss 的方式,数学上完全等价:
Dense KL 是教科书式的方法。它从缓存的 top-100 logits 重建一个完整的 dense teacher-probability 网格,并与 student 自身的 dense log-probability 网格进行比较。这是与在线蒸馏现有做法最接近的版本,因此我们用它作为正确性基线,但它在整个过程中都要在内存中持有完整的词表 × 序列网格,而且是两份。
Forward-chunked KL 保持 teacher 的稀疏性(每个位置只有缓存的 top-100 logits,从不展开成 dense 网格),并逐片计算损失,一次处理一段序列位置。这去掉了 dense teacher 和 dense 比较,而且在我们的基准测试中是三种方法里最快的。不过它仍然有一个盲点:student 自身的 logits,即模型输出层产生的网格,仍然是完整计算并为 backward pass 保留的,因此内存仍会随序列长度急剧增长。
Fused chunked KL,我们的主要贡献,则更进了一步,直接将模型的输出投影融合进 loss 计算。它根本不会产生 student 的完整 logits 网格:它端到端地一次处理一段序列,将该段的隐藏状态投影为 logits,把结果融入正在累积的 loss,然后丢弃该段再进入下一段。Backward pass 是实时重新计算每个段,而不是存储它。代价是该投影要做两次——一次 forward,一次在 backward 中——但作为交换,峰值内存只随序列长度线性增长,而不是随着完整的词表 × 序列大小而飙升。
下面这张 GIF 展示了 dense 和 fused-chunked 两种方法的区别:一种构建完整的比较网格并保留所有内容,另一种则一次构建并丢弃一片,因此内存永远不会超过单个块的大小。
我们已经开源了 chunked-loss 的实现:github.com/CompactifAI/Full-Chunked-KL-Loss

下表将四种设置逐一对比:在线蒸馏,以及刚介绍的三种离线 loss 实现。在单张 H200 GPU 上,以 Llama 3.1 8B Instruct 作为 teacher、3.2B Llama 模型作为 student,上下文 8K token 的条件下进行比较,四种方法的训练 loss 几乎相同,尽管离线运行中每个 token 只用了缓存的 top-100 logits 进行训练。

四种方法的 loss 曲线几乎完全重叠,证实了使用 top-100 缓存 logits 的离线蒸馏相比在线蒸馏是无损的。来源:论文图 2。在这个序列长度下,fused chunked loss 还不是最快的选项,额外的 backward-pass 投影带来了一些速度损失,但它真正的优势只会在上下文长度增长时显现出来,下一节将展示这一点。
为了更鲜明地展示扩展规律,我们在一个玩具输出投影网络(没有 transformer 主干,只有 loss kernel)上做了独立的基准测试。在 32K token 时,峰值内存从 dense loss 的 85.2 GiB 降至完全 chunked 版本的 5.45 GiB,减少了 15.6 倍,而 dense loss 从 64K token 开始就直接 OOM 了。在 256K token 时,完全 chunked loss 使用 11.6 GiB,而次优的 chunked 变体是 134.2 GiB,在这个长度下每次迭代快约 3.3 倍。

在 32,768 token 上下文中蒸馏 GPT-OSS 20B 模型,fused loss 释放的内存让设置从四个 GPU 节点缩减到一个。步时间从 57.0 秒降至 12.23 秒,快了约 5 倍,每张 GPU 的吞吐量从 74.2 升至 345.7 TFLOP/s。
正是这种高效的离线设置使得大规模蒸馏工作首次变得负担得起。得到的紧凑 student,从 Llama 3.1 8B Instruct 蒸馏到约 3.2B 参数,在 BoolQ 和 HellaSwag 上保留了 teacher 的大部分准确性,在 MMLU 上差距在九个百分点以内,而参数量不到一半。

这项工作是 Multiverse Computing 持续研究的一部分,旨在使蒸馏和修复在实际中可以大规模运行,而不仅仅是一次性的配方,让团队能够低成本地迭代。在论文中我们还涵盖了更多的消融实验,例如 loss 函数的选择和序列打包(sequence packing)如何影响恢复质量。
想要查看完整的技术细节,包括 fused chunked loss 背后的闭式梯度以及完整的训练配置?请阅读完整论文,或联系我们的团队,讨论如何将这项技术应用到您自己的蒸馏流程中。
我们已经开源了 chunked-loss 的实现:github.com/CompactifAI/Full-Chunked-KL-Loss