这项技术如何降低知识蒸馏的显存占用?
文章提出两种方法:离线缓存教师模型每个位置最可能的前100个logits,训练时不再同时加载教师模型;以及融合分块KL损失,按序列块端到端计算投影和损失,不保留完整词汇×序列网格,反向传播时按需重算,峰值显存随序列长度线性增长。
(翻译)让知识蒸馏成本足够低以实现大规模运行

A Blog post by Multiverse Computing on Hugging Face
Multiverse Computing 提出两种系统优化,使大模型知识蒸馏可大规模运行:离线缓存教师模型前100个logits,避免训练时同时加载师生模型;融合分块KL散度损失,按块计算并丢弃中间结果,峰值显存显著下降。在32K上下文下显存减少15.6倍,256K上下文仍可运行,并将GPT-OSS 20B蒸馏从四个节点缩减至一个。
文章提出两种方法:离线缓存教师模型每个位置最可能的前100个logits,训练时不再同时加载教师模型;以及融合分块KL损失,按序列块端到端计算投影和损失,不保留完整词汇×序列网格,反向传播时按需重算,峰值显存随序列长度线性增长。
实验显示,四种设置在8K上下文下训练损失几乎重合,说明用缓存的前100个logits做离线蒸馏相对于在线蒸馏几乎无损。
在32K上下文蒸馏GPT-OSS 20B时,显存释放使配置从四个GPU节点缩减到一个,单步时间从57秒降至12.23秒,单卡吞吐从74.2提升至345.7 TFLOP/s。