新算法将知识蒸馏显存降至 6GB 内,支持本地 32K 上下文

ikergarcia1996 · reddit · 2026-08-11

一位开发者开源了针对知识蒸馏中 KL loss 的高效实现方法。通过借鉴 Flash Attention 的分块计算思想,该方法将前向和反向传播进行分块融合,把显存占用从二次方降至线性。

在 32K 上下文长度的设定下,计算 KL loss 的显存需求从原先的约 85GB 暴跌至 5GB 左右,且在长上下文下的计算速度提升了约 3 倍。这打破了以往本地无法进行长上下文知识蒸馏的硬件瓶颈。

该实现需要修改模型的 lm-head 前向传播过程,并支持使用预先缓存的 top-k logits(论文证明使用 top-100 logits 与完整分布效果几乎一致,但大幅减少内存和算力)。作者已同步发布了 GitHub 代码库与 ArXiv 论文。

原文链接 →

「研究」频道最新

更多「研究」频道 AI 资讯 →