大语言模型的高效知识蒸馏:离线Top-K对数概率与融合分块KL损失
Efficient Knowledge Distillation for LLMs: Offline Top-K Logits and a Fused Chunked KL Loss
浏览论文内容
中文总结 AI 辅助
本研究针对大语言模型知识蒸馏成本高的问题,提出离线Top-K对数概率方法与融合分块KL损失,提升训练效率、降低内存占用,相关实现已开源。
中文摘要 AI 辅助
小型语言模型通常是低延迟、低成本和本地部署约束下的唯一选择,但它们很少从头开始训练:压缩模型通常通过知识蒸馏(Knowledge Distillation,KD)恢复。该恢复步骤在很大程度上决定了最终质量,但成本高昂。本文是从业者对如何使蒸馏训练高效的研究,围绕两个系统贡献展开:第一,我们证明离线KD(缓存教师模型的Top-K对数概率一次,让学生模型针对缓存训练)在训练损失与在线蒸馏几乎相同的情况下,将教师模型从内存中移除,每次迭代运行速度提高约29%,在单个H200 GPU上吞吐量最高提高41%;第二,我们引入融合分块KL损失,该损失从不具体化全词汇量的对数概率张量,使峰值内存与序列长度呈线性关系,消除了原本限制上下文长度的内存峰值,让我们在单个GPU上以4倍的上下文长度(32768个token)进行训练。一个仅输出头的玩具基准测试分离了损失核心,证实了其从4K到256K token的内存和迭代速率缩放性。这些共同使大规模修复和数百次 ablation 变得可行,我们还报告了关于损失设计和序列打包的支持性 ablation,并发布了我们的分块损失实现:this https URL。
英文摘要
Small language models are often the only option for deployment under tight latency, cost, and on-premises constraints, but they are rarely trained from scratch: a compressed model is usually recovered through knowledge distillation (KD). This recovery step largely decides the final quality, yet it is expensive. We present a practitioner's study of how to make distillation training efficient, organised around two systems contributions. First, we show that offline KD (caching the teacher's top-$K$ logits once and training the student against the cache) matches online distillation at near-identical training loss while removing the teacher from memory, running about 29\% faster per iteration, and reaching up to 41\% higher throughput on a single H200 GPU. Second, we introduce a \emph{fused, chunked KL loss} that never materialises the full vocabulary-sized logit tensor, making peak memory linear in the sequence length. This removes the memory spike that otherwise caps context length and lets us train at four times the context (32{,}768 tokens) on a single GPU. A separate output-head-only toy benchmark isolates the loss kernel and confirms its memory and iteration-rate scaling from 4K to 256K tokens. Together these make large-scale healing and hundreds of ablations affordable. We also report supporting ablations on loss design and sequence packing. We release our chunked-loss implementation: https://github.com/CompactifAI/Full-Chunked-KL-Loss.