资讯动态

OPD在线策略蒸馏中的反向KL损失:PyTorch实现与梯度陷阱

发布时间:2026/8/29 9:27:43 来源:尧图企业网站定制
训练基于 Transformer 的生成模型时损失函数里一旦出现 KL 这个词很多人会把F.kl_div(input, target)当成普通的交叉熵去调参。可一旦进入 OPD 这类训练分布不断变化的流程KL 的方向选择会直接影响收敛结果。OPD 在很多项目里被写作 On-Policy Distillation在线策略蒸馏它的核心不是让模型去背教师模型的固定答案而是让当前学生模型自己生成样本再用教师模型给这些样本打分。反向 KL 在这里不是一种正则技巧而是和 on-policy 采样天然匹配的损失函数。下面从分布差异讲起用 PyTorch 手写一个可运行的反向 KL 损失并把它接进一个最小的 Transformer OPD 训练循环。1. OPD 是什么为什么它绕不开反向 KL1.1 先约定 OPD 的含义先做一处约定避免不同项目里的缩写冲突。这篇文章里的 OPD 指 Online / On-Policy Distillation即在线策略蒸馏。如果某个代码仓库把 OPD 理解为 Output Probability Distribution或者只把它当成“输出分布之间的某种距离”后面关于 KL 方向、teacher 和 student 的推导依然成立因为最后计算的仍然是两个概率分布之间的差异。Transformer 模型在自回归生成时每个位置都会输出一个词表大小的 logits 向量经过 softmax 之后就是一个概率分布。OPD 的典型流程是学生模型根据当前策略生成一批 token 序列。教师模型对同一序列计算 logits 或概率分布。学生模型在自己生成的序列上重新前向计算与教师分布的损失。反向传播更新学生参数。这和传统的离线蒸馏不同。离线蒸馏通常使用固定语料或教师模型预先离线生成的输出学生没有机会看到自己生成的错误序列。OPD 则是“学生走一步老师纠正一步”训练分布由学生当前策略决定。1.2 离线和在线蒸馏的分布差异离线蒸馏里训练分布基本上是固定的。给定同一个输入教师分布 (p) 和学生分布 (q) 都被约束在同一个 token 位置上。最常用的损失是前向 KL[ D_{KL}(p | q) \sum_i p_i \log \frac{p_i}{q_i} ]这里 (p) 是教师分布(q) 是学生分布。前向 KL 需要从教师分布 (p) 的视角采样也就是重点让 (q) 去覆盖 (p) 认为重要的区域。由于离线语料是固定且充足的这种 KL 方向很自然。但 OPD 的训练分布来自学生自己的采样也就是学生分布 (q)。如果这时候仍然想计算前向 KL严格来说需要估计[ D_{KL}(p | q) \mathbb{E}_{x \sim p} \left[\log \frac{p(x)}{q(x)}\right] ]而实际采样来自 (q)就需要乘上密度比 (p(x)/q(x))。当学生分布和教师分布有明显偏差时这个密度比会出现很大的方差训练不稳定。反向 KL 则不一样[ D_{KL}(q | p) \mathbb{E}_{x \sim q} \left[\log \frac{q(x)}{p(x)}\right] ]采样来自学生自己的 (q)天然就是 on-policy不需要额外的重要性修正。这就是 OPD 选择反向 KL 的数学动机。1.3 OPD 的最小训练循环用一个伪代码可以先建立全局印象repeat until converged: x student_model.generate(prompt) teacher_logits teacher_model(x) student_logits student_model(x) loss reverse_kl(student_logits, teacher_logits) loss.backward() optimizer.step()这个流程里有两个关键点序列 (x) 是学生模型自己生成的不是从固定数据集的标签里抄出来的损失函数使用的是反向 KL而不是把教师输出当作硬标签的交叉熵。下面逐步拆解这两个点。2. 反向 KL 的数学含义和方向陷阱2.1 前向 KL 与反向 KL 的定义假设在某个 token 位置上教师模型给出了分布 (p)学生模型给出了分布 (q)。前向 KL 是[ D_{KL}(p | q) \sum_i p_i \log \frac{p_i}{q_i} ]反向 KL 是[ D_{KL}(q | p) \sum_i q_i \log \frac{q_i}{p_i} ]只看公式很容易觉得“不就是把 (p) 和 (q) 换一下位置吗”。但实际上换位置之后梯度的行为完全不同。前向 KL 对 (q) 的约束是教师概率高的位置学生概率也必须高即使教师分布有几个峰学生也要尽量覆盖所有峰。反向 KL 对 (q) 的约束是学生概率高的位置教师概率不能太低如果学生把概率放到教师不支持的 token 上损失会很大。这种差异在术语上通常表述为前向 KL 是 zero-avoiding反向 KL 是 zero-forcing。前向 KL 倾向于让学生分布覆盖教师分布的支撑集反向 KL 则倾向于让学生分布避开教师分布接近 0 的区域。2.2 前向 KL 与反向 KL 的行为对比对比项前向 KL (D_{KL}(p | q))反向 KL (D_{KL}(q | p))采样分布教师分布 (p)学生分布 (q)强制区域(p) 大的区域(q) 大但 (p) 小的区域典型行为mode covering覆盖多个峰mode seeking收敛到高概率峰风险输出过于分散或平滑输出模式坍缩多样性下降OPD 匹配度需要重要性修正与 on-policy 采样天然匹配在文本生成场景里反向 KL 的 mode seeking 特性并不总是坏事。教师模型通常更强学生模型如果覆盖教师所有可能输出反而容易产生冗余和平滑。反向 KL 会让学生在自己的采样路径上“被老师拉回”输出更集中、更贴近老师认可的高概率区域。2.3 Transformer 输出层如何识别 KL 方向Transformer 的注意力层和 FFN 层不会直接参与 KL 计算。模型前向传播走到最后一层后会得到形状为[batch_size, seq_len, vocab_size]的 logits。对 logits 做log_softmax就能得到每个位置的概率对数这就是 KL 的输入。在代码里最容易出错的地方是F.kl_div(input, target)的第一个参数必须是“对数概率”第二个参数是“概率或者对数概率”。很多人的input和target写反大概率是没理解清楚当前算的是前向还是反向。推荐先用下面这个明确方向前向 KLteacher 是 targetstudent 是拟合方。反向 KLstudent 的分布是采样方teacher 的分布是参考分布。3. 在 PyTorch 里手撕反向 KL3.1 直接按定义实现先写一个最纯粹的反向 KL。student_logits来自学生模型teacher_logits来自教师模型。为了避免数值不稳定不要先softmax再取log直接用log_softmax。import torch import torch.nn.functional as F def reverse_kl_from_logits(student_logits, teacher_logits): log_q torch.log_softmax(student_logits, dim-1) log_p torch.log_softmax(teacher_logits, dim-1) q torch.exp(log_q) per_token torch.sum(q * (log_q - log_p), dim-1) return per_token.mean()这个函数有几个细节log_q - log_p等价于 (\log(q_i/p_i))。q由log_q还原出来保持学生分布是计算图的一部分梯度能通过q流回学生模型。最后的per_token.mean()是对所有 token 位置求平均而不是求和。求和会让长序列的 loss 天然比短序列大。3.2 用 F.kl_div 实现并防止参数传反PyTorch 的F.kl_div默认计算[ \sum_i target_i \times (\log target_i - input_i) ]因此如果target是概率input必须是 log 概率。用这个 API 实现前向 KL 和反向 KL 时参数方向完全不同。# 前向 KL: KL(P || Q) forward_kl F.kl_div( inputtorch.log_softmax(student_logits, dim-1), # log Q targettorch.softmax(teacher_logits, dim-1), # P reductionbatchmean ) # 反向 KL: KL(Q || P) reverse_kl F.kl_div( inputtorch.log_softmax(teacher_logits, dim-1), # log P targettorch.softmax(student_logits, dim-1), # Q reductionbatchmean )注意反向 KL 里target是学生分布 (q)并且q必须来自可导的学生 logits。教师模型的 logits 应该被torch.no_grad()包住否则梯度会错误地流到教师模型。想计算的分布F.kl_div的 inputF.kl_div的 targetlog_target 参数前向 KL (P | Q)(\log Q)(P)False反向 KL (Q | P)(\log P)(Q)False前向 KLtarget 给 log 概率(\log Q)(\log P)True反向 KLtarget 给 log 概率(\log P)(\log Q)True3.3 加入 temperature 和 padding mask实际训练中很少直接用裸 logits。教师模型输出可能过于尖锐也会有很多 padding token。先用 temperature 对分布做平滑再通过 mask 屏蔽 padding 或 prompt 位置。def reverse_kl_with_temperature_and_mask( student_logits, teacher_logits, token_mask, temperature1.0, ): student_logits student_logits / temperature teacher_logits teacher_logits / temperature log_q torch.log_softmax(student_logits, dim-1) log_p torch.log_softmax(teacher_logits, dim-1) q torch.exp(log_q) per_token torch.sum(q * (log_q - log_p), dim-1) token_mask token_mask.float() denom token_mask.sum().clamp(min1.0) return (per_token * token_mask).sum() / denom * (temperature ** 2)这里乘上temperature ** 2是因为 Hinton 蒸馏里常见的尺度约定温度升高后软目标更平滑梯度的绝对值会下降乘以 (T^2) 可以补偿这种下降。这个乘法不会改变 KL 方向只影响整体 loss 的尺度。3.4 计算图切分哪些地方要 detach手写反向 KL 时最大的计算图陷阱是.detach()用错位置。教师模型的log_p必须 detach。如果教师模型参数不更新可以直接把教师模型前向包在torch.no_grad()里。学生模型的log_q和q都不能 detach。反向 KL 里学生分布 (q) 既是期望的采样分布也是被优化的概率分布。直接计算形式的 loss 中q * (log_q - log_p)的梯度既来自括号里的log_q也来自外层的

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价