让知识蒸馏的成本足够低,以便大规模运行
Hugging Face · · 发布于 2026-08-10 · 32 分钟阅读
本文提出两种系统改进——离线缓存教师模型的Top-K logits和融合分块KL损失,将知识蒸馏的GPU内存需求从约250GB降至约128GB,使其可在单张GPU上运行。
标准蒸馏恢复成本高
在线KL蒸馏需要同时加载教师和学生,每个token产生全词汇表概率分布。以gpt-oss-120b为例,序列长度32K、batch size 4时,教师概率张量约50GB,单次迭代峰值显存约250GB,超过单块H200或B200的容量。
两项系统改动实现降本
一是离线缓存教师模型每个位置的top-100 logits,教师无需在训练时驻留显存;二是融合分块KL损失,将输出投影与损失计算融合,按序列分块处理,避免物化全词汇表×序列长度的矩阵。
精度保持与显存削减
在8K上下文单H200上,在线蒸馏、离线dense、前向分块和融合分块四种方法的训练损失曲线几乎重叠;融合分块KL的峰值显存从在线蒸馏的102.8GB降至58.3GB。
长上下文扩展能力
隔离基准中,32K上下文下峰值显存从85.2GiB降至5.45GiB,削减15.6倍;密集损失在64K起失败,而融合分块在256K时仅用11.6GiB,且比次优分块变体快约3.3倍。
实际蒸馏案例
对GPT-OSS 20B模型在32,768上下文下蒸馏,所需GPU节点从4个缩减到1个,步时间从57.0秒降至12.23秒,约5倍提升,每GPU吞吐从74.2升至345.7 TFLOP/s。
学生模型表现
从Llama 3.1 8B Instruct蒸馏到约3.2B参数,学生保留BoolQ和HellaSwag上的大部分准确率,MMLU差距约9分,参数不到教师一半。
SOURCES
共 1 个来源来源与引用
官方原文Hugging Face 原始发布2026-08-10 · 一手信息https://huggingface.co/blog/MultiverseComputingCAI/efficient-knowledge-distillation