Hugging Face 提出高效知识蒸馏方案:单卡即可蒸馏千亿模型
Hugging Face 论文提出离线 Top-K 缓存与融合分块 KL 损失两项改造,将千亿模型蒸馏显存峰值从约 25…
Hugging Face 在其官方博客上介绍了最新论文《Efficient Knowledge Distillation for LLMs: Offline Top-K Logits and a Fused Chunked KL Loss》,提出两项系统级改造,将大语言模型知识蒸馏的显存与计算开销大幅压缩。论文给出的数据显示,针对 gpt-oss-120b 的蒸馏训练峰值显存从约 250GB 降至约 128GB,首次使超长上下文(32K+)的蒸馏恢复训练可在单张 GPU 上完成,也让大规模蒸馏实验不再依赖数百卡集群。
蒸馏恢复为何昂贵
知识蒸馏是当前开源大模型压缩的主流做法。OpenAI 的 gpt-oss、阿里通义 Qwen、智谱 GLM、月之暗面 Kimi 等最新一代开源权重模型参数规模持续攀升,其中 Kimi-K3 拥有 2.8 万亿参数,仅加载就需约 3TB 显存。NVIDIA Nemotron 3 Puzzle 75B、Multiverse Computing Hypernova 60B 等近期发布的高质量压缩模型也都依赖蒸馏来恢复能力。
传统在线蒸馏流程需要教师模型与学生模型同时驻留显存:每一步训练中,教师都要重跑一次完整前向,生成覆盖整个词表的概率分布,并与学生的分布计算 KL 散度。以 gpt-oss-120b 为例,其词表大小为 201,088,在序列长度 32K、batch size 4 的设置下,仅教师概率张量的形状就达到 4 × 201,088 × 32,768,bfloat16 精度下约占 50GB 显存;叠加梯度、激活、权重与优化器状态后,单次迭代峰值显存可达约 250GB,超过单张 H200(141GB)甚至 B200 的承载能力。
两项系统级改造
论文的核心贡献由两个相互独立又彼此互补的改动构成。
离线 Top-K 蒸馏
传统做法下教师模型每个训练步都要重新前向一次,但其实教师在整个训练过程中并不会改变。论文的思路是:预先对每个 token 位置计算一次教师输出,仅缓存概率最高的 Top-100 logits,后续训练不再需要把教师加载进显存,同一份缓存还可在多次消融实验中反复复用。
融合分块 KL 损失
KL 损失本身的开销同样巨大:它需要为词表中每个词、序列中每个位置各计算一个数值,本质上是一个词表大小 × 序列长度的巨大矩阵,默认实现需要先把整张表全部物化才能开始求和。论文对比了三种数学上等价但工程效率迥异的实现方式:
- Dense KL:教科书做法,将教师 Top-100 重建为稠密网格与学生稠密 log 概率对比,作为正确性基线,但需要同时持有两份完整词表 × 序列的矩阵。
- Forward-chunked KL:保持教师稀疏(仅 Top-100),按序列分片逐块计算损失,速度最快;但学生的完整 logits 仍需在反向传播前保留,显存随序列长度仍快速增长。
- Fused chunked KL:论文的核心方案,将模型输出投影直接融合进损失计算,从不构造学生的完整 logits 网格,而是按序列分块端到端处理:先投影隐藏状态到该块的 logits、立即折算进累积损失、然后丢弃,再处理下一块;反向时按需即时重算,峰值显存仅随序列长度线性增长,不再出现「词表 × 序列」的尖峰。
代价是该投影在前向与反向中各执行一次,但换来的显存曲线几乎平直。
效果与意义
论文给出的对比图显示,Dense KL 在蒸馏 gpt-oss-120b 时显存峰值约 250GB,超过单卡 H200 的容量;而融合分块 KL 损失全程不超过约 128GB。两者结合意味着:教师只需离线计算一次 Top-K 缓存,学生训练阶段不再依赖教师驻留显存,损失计算又避免了词表 × 序列矩阵的尖峰。
实际意义在于两点。其一,长上下文(32K 及以上)的蒸馏恢复首次具备单卡可行性,对算力有限的中小团队尤为关键;其二,单次蒸馏成本下降后,团队可以在更大规模上做超参与数据消融,加速压缩模型迭代。该工作已以论文形式发布,配套实现计划后续开源,对当前以蒸馏为主要压缩路径的开源大模型生态具有直接参考价值。
