
图片:METAL
摘要
- Multiverse Computing公司公开了两项系统改进,可大幅降低知识蒸馏训练中的显存占用
- 通过预先缓存教师模型的top-K logits,消除了需要与学生模型同时载入内存的问题
- 应用了不生成完整词表×序列长度矩阵的新型KL散度损失函数
Multiverse Computing公司通过Hugging Face博客公开了一篇论文,内容涉及降低大型语言模型知识蒸馏(knowledge distillation)训练成本的两项系统改进。论文标题为《Efficient Knowledge Distillation for LLMs: Offline Top-K Logits and a Fused Chunked KL Loss》。
知识蒸馏是让较小的学生(student)模型学习模仿大型教师(teacher)模型性能的技术。传统方法需要在训练的每一步都将教师模型同时载入内存以计算整个词表的概率分布,因此显存负担很大。例如,词表规模为201,088的gpt-oss-120b模型,在序列长度32K、批大小为4的条件下,仅教师概率张量在bf16精度下就约达50GB,若再加上梯度、激活值和优化器状态等,训练单步的显存占用最高可能飙升至约250GB。
两项改进
研究团队通过预先一次性计算并缓存教师模型的top-K logits,消除了训练过程中需要将教师模型与学生模型同时保留在内存中的必要性。此外,团队还引入了一种新型的内存高效KL散度损失函数,该函数不会直接生成词表规模×序列长度大小的矩阵,相比PyTorch或NVIDIA Megatron-Bridge等库的默认实现,大幅降低了显存占用。
研究团队表示,将这两项改进结合起来,仅用一张GPU就能实现长上下文(long-context)恢复训练,大规模实验的运行成本也能实质性降低。这一成果的发布正值超大规模模型不断涌现之际——例如参数规模达2.8万亿、仅加载就需约3TB显存的Kimi-K3,与此同时,NVIDIA的Nemotron 3 Puzzle 75B、以及Multiverse Computing自家的Hypernova 60B等压缩模型也相继问世。





评论