知识蒸馏(KD)中的KL#
Geoffrey Hinton 在其 2015 年的论文Distilling the Knowledge in a Neural Network中使用了KL散度进行判别模型度知识蒸馏,形成了Supervised KL Distillation的经典范式。其训练目标函数可以写作
L=αLhard+βLsoft=−(1−λ)E(x,y)∼(X,Y)ylogqθ(1)(y∣x)+λT2Ex∼XDKL(p(T)(⋅∣x)∥qθ(T)(⋅∣x))=−(1−λ)E(x,y)∼(X,Y)ylogqθ(1)(y∣x)+λT2Ex∼X,y∗∼p(T)(⋅∣x)logqθ(T)(y∗∣x)p(T)(y∗∣x)其中 p(T) 和 qθ(T) 分别代表教师模型和学生模型在温度 T 下的概率。公式中对于简单的判别任务,直接在所有类别上计算KL散度,可以得到精确的数值。
GKD: On-Pollicy LLM 蒸馏#
论文 On-Policy Distillation of Language Models: Learning from Self-Generated Mistakes将上述方法直接迁移到LLM,将传统的监督信号和LLM强化学习的On-policy生成相结合,提出了GKD(Gererative Knowledge Distillation)的训练范式。其训练目标函数可以写作
L=(1−λ)E(x,y)∼(X,Y)∣y∣1n=1∑∣y∣DKL(p(⋅∣y<n,x)∥qθ(⋅∣y<n,x))+λEx∼X,y∗∼qθ(⋅∣x)∣y∗∣1n=1∑∣y∗∣DKL(p(⋅∣y<n∗,x)∥qθ(⋅∣y<n∗,x))=(1−λ)E(x,y)∼(X,Y)∣y∣1n=1∑∣y∣Ev∼p(⋅∣y<n,x)logqθ(v∣y<n,x)p(v∣y<n,x)+λEx∼X,y∗∼qθ(⋅∣x)∣y∗∣1n=1∑∣y∗∣Ev∼p(⋅∣y<n,x)logqθ(v∣y<n∗,x)p(v∣y<n∗,x)当 λ=1 时即为 On-Policy GKD:
L=Ex∼X,y∼p(⋅∣x)∣y∣1n=1∑∣y∣DKL(p(⋅∣y<n,x)∥qθ(⋅∣y<n,x))=Ex∼X,y∼p(⋅∣x)∣y∣1n=1∑∣y∣Ev∼p(⋅∣y<n,x)logqθ(v∣y<n,x)p(v∣y<n,x)实际应用中,为了加快计算,常使用 Top-K 概率截断来加速计算,同时使用 K3 估计来近似计算:
L=Ex∼X,y∼p(⋅∣x)∣y∣1n=1∑∣y∣DKL(p(⋅∣y<n,x)∥qθ(⋅∣y<n,x))≈Ex∼X,y∼p(⋅∣x)∣y∣1n=1∑∣y∣v∈Vtop-k∑(qθ(v∣y<n,x)−p(v∣y<n,x)+p(v∣y<n,x)logqθ(v∣y<n,x)p(v∣y<n,x))K3 估计很好地保证了在 Top-k 截断后 KL 估算的非负性。
MiniLLM: 反向 KL + 蒙特卡洛采样实现高效蒸馏#
论文 MiniLLM: Knowledge Distillation of Large Language Models提出了使用Reverse KL进行模型蒸馏。其训练目标为最小化序列级别分布的Reverse KL,可以写作
L=Ex∼X,y∼qθ(⋅∣x)logp(y∣x)qθ(y∣x)其中qθ为学生模型,p为教师模型。
考虑其梯度:
∇L=∇Ex∼X,y∼qθ(⋅∣x)logp(y∣x)qθ(y∣x)=Ex∼Xy∑∇(qθ(y∣x)logp(y∣x)qθ(y∣x))=Ex∼Xy∑logp(y∣x)qθ(y∣x)∇qθ(y∣x)+Ex∼Xy∑qθ(y∣x)∇logqθ(y∣x)=Ex∼Xy∑qθ(y∣x)logp(y∣x)qθ(y∣x)∇logqθ(y∣x)+Ex∼Xy∑qθ(y∣x)∇logqθ(y∣x)=Ex∼X,y∼qθ(⋅∣x)(logp(y∣x)qθ(y∣x)+1)∇logqθ(y∣x)=Ex∼X,y∼qθ(⋅∣x)n=1∑∣y∣m=1∑∣y∣logp(ym∣y<m,x)qθ(ym∣y<m,x)+1∇logqθ(yn∣y<n,x)=Ex∼X,y∼qθ(⋅∣x)n=1∑∣y∣m=n∑∣y∣logp(ym∣y<m,x)qθ(ym∣y<m,x)+1∇logqθ(yn∣y<n,x)可以发现梯度可以写成策略梯度的形式。因此可以使用强化学习的方法,通过蒙特卡洛采样的方式估计梯度。
Reverse KL的主要优势如下图所示,学生模型不会尝试拟合教师的完整分布,而是去拟合和学生模型分布相近的一个子分布。该特点使得蒸馏后学生模型不容易损失原有的能力。

蒙特卡洛采样会引入极大的方差,在整个 sequence 级别计算 KL 还会使方差累积。因此,当前更常见的方法是仅优化单步的 KL 损失,即
L=Ex∼X,y∼qθ(⋅∣x)n=1∑∣y∣Eyn∗∼qθ(⋅∣y<n,x)logp(yn∗∣y<n,x)qθ(yn∗∣y<n,x)其梯度为
∇L=n⩾1∑Ex∼X,y∼qθ(⋅∣x)Eyn∗∼qθ(⋅∣y<n,x)1{∣y∣⩾n}(logp(yn∗∣y<n,x)qθ(yn∗∣y<n,x)−1)∇logqθ(yn∗∣y<n,x)=n⩾1∑Ex∼X,y<n∼qθ(⋅∣x)Eyn∗∼qθ(⋅∣y<n,x)1{∣y<n∣=n−1}(logp(yn∗∣y<n,x)qθ(yn∗∣y<n,x)−1)∇logqθ(yn∗∣y<n,x)=n⩾1∑Ex∼X,y⩽n∼qθ(⋅∣x)1{∣y⩽n∣=n}(logp(yn∣y<n,x)qθ(yn∣y<n,x)−1)∇logqθ(yn∣y<n,x)=n⩾1∑Ex∼X,y∼qθ(⋅∣x)1{∣y∣⩾n}(logp(yn∣y<n,x)qθ(yn∣y<n,x)−1)∇logqθ(yn∣y<n,x)=Ex∼X,y∼qθ(⋅∣x)n=1∑∣y∣(logp(yn∣y<n,x)qθ(yn∣y<n,x)−1)∇logqθ(yn∣y<n,x)=Ex∼X,y∼qθ(⋅∣x)n=1∑∣y∣logp(yn∣y<n,x)qθ(yn∣y<n,x)∇logqθ(yn∣y<n,x)因此可以直接将 logp(yn∣y<n,x)qθ(yn∣y<n,x) 作为强化学习的 Advantage 进行训练。
DeepSeek-V4: 全词表精确反向 KL 计算#
DeepSeek-V4的技术报告 DeepSeek-V4: Towards Highly Efficient Million-Token Context Intelligence 中提到了 DeepSeek-V4 在全词表上计算反向 KL 以实现精确的数值计算,即
L=Ex∼X,y∼p(⋅∣x)∣y∣1n=1∑∣y∣DKL(qθ(⋅∣y<n,x)∥p(⋅∣y<n,x))=Ex∼X,y∼p(⋅∣x)∣y∣1n=1∑∣y∣v∈V∑qθ(v∣y<n,x)logp(v∣y<n,x)qθ(v∣y<n,x)这在超大参数量、多教师模型以及超长上下文的推理场景下造成了极大的性能挑战。为了解决这个问题,DeepSeek-V4 在 OPD 训练阶段将教师权重全部卸载至集中式分布式存储中,并在教师前向传播期间按需加载,采用类似 ZeRO 的参数分片技术以缓解 I/O 与 DRAM 压力。此外,在教师模型前向传播中缓存最后一层 hidden state 并在训练中重新复原 logits,避免显式实例化 logits 带来极大的内存负担。