反向KL在Transformer输出分布中的陷阱与正确选型

发布时间:2026/8/30 5:23:44
反向KL在Transformer输出分布中的陷阱与正确选型 手写过 Transformer decode 过程的人大概率都见过这个瞬间最后一层线性层吐出一堆 logitssoftmax 之后得到一个概率分布然后你拿着这个分布去算交叉熵、算 KL 损失。大多数时候我们只关心 loss 是不是在下降但这个输出概率分布Output Probability Distribution后面我统一简写成 OPD里藏着一个特别容易翻车的细节——反向 KL。我第一次真正注意到反向 KL是在一次知识蒸馏实验里。教师模型和学生模型都是标准 Transformer decoder蒸馏损失把教师和学生的 OPD 拿来做 KL 对齐。loss 曲线一路下降看起来一切正常但验证集上的生成结果却越来越“稳”稳到几乎每个句子都缩成那几条高频路径。后来排查发现我把 KL 的方向写反了。方向错误的 loss 依然能下降因为它也是一个有效的散度但优化目标已经和“让学生学教师的完整分布”完全不同了。这个经历让我意识到Transformer 的细节并不只在 attention mask、位置编码、层归一化这些“大件”上OPD 上的 KL 方向同样是一个能决定模型行为的隐藏旋钮。反向 KLReverse KL也叫逆 KL与前向 KLForward KL看起来只是交换了公式里的两个分布实际行为却差得很远。这篇笔记就把这个细节拆开OPD 是什么反向 KL 在其中怎么算、怎么用、怎么踩坑以及真正落地时应该怎么选方向。1. 先搞清楚 OPD 和反向 KL 到底在说什么1.1 从 logits 到 OPD一个概率分布是如何诞生的在 Transformer 里无论是 decoder 的下一 token 预测还是 encoder 的分类输出模型最后通常都会产生一组 logits。它的形状一般是[batch_size, seq_len, vocab_size]对这个张量在最后一维做 softmax就得到 OPD。这一步看起来简单但它决定了模型对每个位置“认为哪个 token 更可能”。这里有两个细节值得手撕时记清楚logits 的数值范围没有约束可能是正也可能为负甚至出现几十几百的极端值。softmax 之后每行概率总和为 1但分布的尖锐程度完全由 logits 的尺度决定。如果在 softmax 之前给 logits 除以一个温度T当T 1时分布会更平滑T 1时分布会更尖锐。温度缩放是后续 KL 计算里最常见的调节手段因为它直接改变 OPD 的形状而 KL 对分布形状非常敏感。在 PyTorch 里最小结构是这样import torch import torch.nn.functional as F batch_size, seq_len, vocab_size 2, 8, 32 logits torch.randn(batch_size, seq_len, vocab_size) probs F.softmax(logits, dim-1) log_probs F.log_softmax(logits, dim-1) print(probs.sum(dim-1)) # 每行应该接近 1 print(log_probs.shape) # [batch, seq, vocab]log_probs不是简单对probs取log而是直接由log_softmax计算。原因后面会专门说但这是实现 KL 的第一步。1.2 反向 KL 的定义两个分布之间的方向性距离在 OPD 上定义两个分布教师分布P和学生分布Q。KL 散度有两个常见方向前向 KL: KL(P || Q) sum_x P(x) * (log P(x) - log Q(x)) 反向 KL: KL(Q || P) sum_x Q(x) * (log Q(x) - log P(x))注意这里的“前向”和“反向”不是绝对标准不同资料里的叫法可能相反。关键在于哪个分布出现在加权的位置。前向 KL 用P作为权重所以P概率高的区域贡献大。如果P在某处有概率但Q几乎为 0那么log Q(x)会变成很大的负值导致散度爆炸。前向 KL 会“逼着”Q覆盖P的整个支撑区域。反向 KL 用Q作为权重所以Q自己低概率区域的影响很小。即使P在某个地方有高概率只要Q不认为自己会走到那里损失也不关心。反向 KL 更偏向“让学生安分守己地在自己选中的高概率区域接近教师”。用一句话概括前向 KL 鼓励“覆盖”反向 KL 容易“保守”。在 OPD 上这意味着反向 KL 可能会让学生在生成时主动避开教师分布中那些尾部 token哪怕教师确实给了它们概率。1.3 为什么 Transformer 细节里到处都有 KLKL 散度在 Transformer 训练中出现频率极高最常见的地方包括知识蒸馏让学生 OPD 去匹配教师 OPD。生成约束防止模型输出过平滑或过于尖锐的正则项。变分上下文在潜在变量上做 KL 正则。批量校准在输出层做分布校准。从数学上看交叉熵和 KL 有直接关系CrossEntropy(P, Q) H(P) KL(P || Q)其中H(P)是P的熵。如果P是固定标签分布那么最小化交叉熵等价于最小化前向 KL。这也是为什么分类任务里大量接触的都是前向 KL反向 KL 相对陌生。但在 distillation、强化学习式策略优化、生成分布约束等场景里反向 KL 会突然冒出来。如果不理解它的行为就很容易出现开头说的那种“loss 在降模型却在变怪”的情况。2. 反向 KL 在 Transformer 场景中的三个典型使用位置2.1 知识蒸馏教师 OPD 和学生 OPD 之间该选哪个方向知识蒸馏是最容易暴露 KL 方向问题的场景。假设教师分布是P学生分布是Q通常我们想要的是让学生学教师的分布。很多开源代码里直接用loss F.kl_div(student_log_probs, teacher_probs, reductionbatchmean)这里的input是学生 log 概率target是教师概率所以 PyTorch 实际计算的是KL(teacher || student)也就是前向 KL。它要求学生尽可能覆盖教师的所有高概率区域。如果把它改成反向 KLKL(student || teacher)代码需要写成loss F.kl_div(teacher_log_probs, student_probs, reductionbatchmean)这时加权的是学生自身的概率分布。学生会更关注“自己大概率选哪个 token这个 token 在教师那里是不是也够高”。它不会主动去覆盖教师分布里的低概率 token。从我的实验经验看蒸馏任务默认使用前向 KL 更稳妥因为教学目标是让学生复现教师的整体分布。反向 KL 如果单独使用容易让学生的生成分布变窄尤其当学生模型容量不足或训练轮次不够时它会选择只学习教师的高置信路径牺牲多样性。它更适合那些希望学生“生成时要稳不要乱发散”的场景。2.2 序列生成约束防止 OPD 过平滑或过尖锐在自回归生成任务里模型输出的 OPD 很大程度上决定了文本风格和重复程度。一个常见问题是过平滑模型对每个位置都给出接近均匀的概率生成结果泛泛而谈。另一个常见问题是过尖锐某个 token 概率几乎为 1生成结果变成复读机。反向 KL 可以作为正则项加入训练 loss。因为它由学生分布Q加权所以它会惩罚“学生自信但教师不认可”的 token。如果学生的Q在某个 token 上概率很高而教师的log P很小那么反向 KL 会产生较大损失从而抑制学生的盲目自信。这有点像把学生拉回到“你已经选定的路径”上确保这条路在教师视角里是靠谱的。但它不会主动帮学生探索教师认可的其他路径所以如果你希望生成结果保持多样性不要只依赖反向 KL还要搭配前向 KL 或多样性奖励。2.3 变分和中间层对齐反向 KL 作为正则项的边界在变分 Transformer 或 latent variable model 中反向 KL 更常出现在对潜在变量z的优化上KL(q(z) || p(z))。这里q(z)是近似后验p(z)是先验。反向 KL 会让q(z)只在它认为合理的区域里贴近先验遇到先验低概率区会自动避开。但要特别说明这种用法面对的是潜在变量分布和本文讨论的 OPD 不完全一样。ODP 是模型输出的 token 概率分布属于模型结果层。在结果层使用反向 KL更多是直接约束生成行为而不是变分推断。实际工程里不要把这两类混为一谈否则你会发现 loss 的公式和梯度行为都对不上。3. 手撕实现在 Transformer 里计算反向 KL 的完整细节3.1 最小实现自己写一次 KL而不是只调用黑盒很多初学者直接调用F.kl_div但一旦遇到方向、温度、mask 问题就很难排查。建议先手写一次理解公式对应的张量操作。下面这段代码实现一个通用的 KL 函数可以指定方向import torch import torch.nn.functional as F def kl_divergence( p_logits, q_logits, temperature1.0, use_reverseFalse, maskNone, ): # p_logits: 参考分布 logits比如教师 # q_logits: 需要优化的分布 logits比如学生 # 默认计算 KL(p || q) if use_reverse: # 反向 KL: KL(q || p) reference_logits, target_logits q_logits, p_logits else: # 前向 KL: KL(p || q) reference_logits, target_logits p_logits, q_logits p_log_probs F.log_softmax(reference_logits / temperature, dim-1) q_log_probs F.log_softmax(target_logits / temperature, dim-1) p_probs torch.exp(p_log_probs) # 散度值 shape: [batch, seq] kl (p_probs * (p_log_probs - q_log_probs)).sum(dim-1) if mask is not None: kl kl * mask.float() valid_count mask.float().sum() return kl.sum() / valid_count.clamp_min(1.0) else: return kl.mean()这里的关键是p_log_probs实际来自哪个分布决定了谁在加权。3.2 数值稳定性不要用 softmax 后的概率去取 log计算 KL 时标准做法是在 log 空间完成先log_softmax再用概率乘两个 log 概率之差。千万不要这样写p_probs F.softmax(p_logits, dim-1) q_probs F.softmax(q_logits, dim-1) kl (p_probs * (torch.log(p_probs) - torch.log(q_probs))).sum(dim-1)这样做有两个问题log(p_probs)在p_probs为 0 时会出现-inf反向传播时梯度容易变成 NaN。softmax后得到的是一个经过了归一化的浮点张量再取对数会累积舍入误差尤其当某一项概率极小但并非 0 时误差会被放大。log_softmax会在底层使用log-sum-exp技巧数值上稳定得多。所以在任何 Transformer 代码里只要涉及 KL 或交叉熵都优先使用log_softmax的输出。3.3 一个可运行的蒸馏小片段下面这个例子模拟教师和学生各有一批 logits然后计算反向 KL 损失并加上温度控制。import torch import torch.nn.functional as F torch.manual_seed(0) batch_size, seq_len, vocab_size 2, 16, 64 teacher_logits torch.randn(batch_size, seq_len, vocab_size) student_logits torch.randn(batch_size, seq_len, vocab_size) # padding mask最后 4 个位置是 padding mask torch.ones(batch_size, seq_len) mask[:, -4:] 0 # 计算反向 KL loss kl_divergence( teacher_logits, student_logits, temperature2.0, use_reverseTrue, maskmask, ) print(loss.item())在这个实现里temperature 会把 logits 缩小分布变得平滑反向 KL 的“保守”效应会被弱化。你可以观察不同temperature下 loss 的变化这有助于理解温度如何影响优化目标。3.4 温度缩放它会改变反向 KL 的行为温度对反向 KL 的影响不是线性的。高温会让两个分布都趋于均匀KL 值整体下降梯度也会变小低温会让分布更尖锐反向 KL 对学生的高置信 token 更敏感梯度可能更大但也更容易不稳定。一个常见做法是先固定T1.0跑通流程确认 mask、loss 方向都正确再开始调温度。如果训练初期 loss 就出现 NaN优先检查是否因为温度过低导致 logits 出现极端值或者因为教师分布中存在零概率。4. 最容易踩坑的五个细节和排查链路4.1 方向写反loss 正常行为不同这是最隐蔽的问题。因为KL(P||Q)和KL(Q||P)都是非负的都能作为 loss 优化所以反向传播时 loss 会下降。但梯度引导的方向完全不同。我在文章开头提到的实验就是这种典型把F.kl_div的两个参数顺序搞反loss 曲线照样平滑但 student 的生成分布越来越窄。原因是反向 KL 让学生只需要在自己已经高概率的区域靠近教师教师分布里的长尾部分全部被忽略。排查方法在训练日志里同时记录两个方向的 KL 值比如KL_forward和KL_reverse。如果只有一个方向在下降另一个长期不变你就需要停下来确认优化目标到底是什么。4.2 把概率当 log 概率使用F.kl_div的输入签名是input作为 log probabilitiestarget作为 probabilities。如果误把概率传入input损失值会非常奇怪甚至为负。更常见的是自己手写公式时先对概率取log再参与运算。这通常是数值不稳定的根源。建议在实现阶段写一个断言确保传入log_prob的张量最大值小于等于 0assert log_probs.max().item() 1e-6, log_probs 应该是对数概率4.3 没有处理 padding token在 Transformer 的自回归训练中序列会 padding 到相同长度。如果不对 padding 位置做 maskKL 损失会被这些无意义位置主导。比如一个 batch 里有两条序列一条长度为 16另一条长度为 8padding 到 16。如果直接对整张[batch, seq, vocab]求均值短序列的 padding 位置也会参与计算它们对应的 logits 是随机初始化的产生的 loss 会稀释真正 token 的梯度。正确做法是维护一个mask有效位置为 1padding 位置为 0。计算损失时对 mask 后的值求和或求平均只统计有效位置。4.4 反向 KL 遇到零概率导致 NaN反向 KL 的公式是KL(Q || P) sum_x Q(x) * (log Q(x) - log P(x))如果某个 token 在Q中概率为 0通常约定0 * log 0 0问题不大。但如果Q(x) 0而P(x) 0那么log P(x) -inf反向 KL 直接变成无穷大。这在 OPD 蒸馏里很常见教师分布经过了 softmax虽然概率很小但一般不会严格为 0不过如果使用了硬标签、或对教师输出做了 top-k 截断就可能存在严格 0 概率的区域。解决办法包括对教师 logits 做温度缩放使概率更平滑。对P(x)加一个极小的光滑项例如先P P 1e-8再归一化。在计算 KL 时对log_P做clamp_min(-100)避免直接无穷大。4.5 排查链路从张量到损失值的五个检查点当反向 KL 出现异常时我一般按这个顺序排查检查点观察现象可能原因logits 张量出现 NaN 或 Inf上游梯度爆炸、学习率过大softmax 求和每行不等于 1在错误维度 softmaxlog_probs最大值小于等于 0误用torch.log(probs)mask 统计有效 token 数量正确mask 类型不是 float或维度和 logits 不一致loss 方向前向/反向记录异常F.kl_div参数顺序写反这五层排查基本能覆盖 90% 的 KL 问题。如果仍然异常就要检查教师和学生是否共享参数、梯度是否正常传导以及 teacher 是否会因为no_grad导致梯度断流。5. 怎么选方向一个判断框架和落地建议5.1 四步选型框架面对一个具体任务时可以按下面四步来决定用哪个方向的 KL明确哪个分布是“参考标准”哪个分布是“被优化对象”。确认你希望被优化对象“覆盖参考标准”还是“在自身高概率区域贴近参考标准”。如果任务要求多样性和完整覆盖优先前向 KL如果任务要求稳定和避免发散可以考虑反向 KL但必须加监控。设置温度后在小验证集上同时观察两个方向的 loss 和实际生成质量而不是只看训练 loss。可以用这个简化的判断表目标推荐方向理由学生完整学习教师分布前向 KL用教师概率加权覆盖教师高概率区域学生只生成高置信 token反向 KL用学生概率加权忽略低概率尾部防止学生过平滑反向 KL惩罚“学生自信但教师不认可”的区域保持生成多样性前向 KL 或组合反向 KL 容易收敛到少数模式5.2 反向 KL 不适合什么场景反向 KL 不是万能。以下情况我会尽量避免单独使用它多标签输出或高熵任务如果正确分布本身有很多可能 token反向 KL 会让学生丢弃尾部可能导致覆盖面不足。教师分布有大量零概率区域数值不稳定容易出现 NaN。训练初期模型分布非常不均匀学生刚开始学OPD 可能集中到错误 token反向 KL 会放大这种错误自信让模型一直卡在局部最优。需要生成结果具备创造性的场景反向 KL 的保守性会成为限制。在这些场景里更好的做法是使用前向 KL 和反向 KL 的加权组合total_kl alpha * KL(P || Q) beta * KL(Q || P)但组合后需要调节两个系数工程成本会增加。我的建议是先用单一方向把基线跑出来再决定要不要组合。5.3 长期使用需要补上的工程化能力如果你的项目要长期使用反向 KL只实现一个 loss 函数是不够的。还需要补上日志记录单独记录KL_forward、KL_reverse、KL_combined的值监控趋势。分布可视化从验证集里抽取几个样本打印教师和学生 OPD 中 top-10 token观察差异。梯度检查确认反向 KL 对 student logits 的梯度在学校自己高概率 token 上更大而不是只对某个维度敏感。checkpoint 选择不能只看 distillation loss还要结合下游指标选 checkpoint因为反向 KL 的 loss 降低可能伴随多样性下降。这些工程化能力其实和模型结构无关但决定了你能不能让“细节”变成一个可复用的能力。5.4 一个可复用的检查清单使用反向 KL 之前建议逐项确认我已经明确哪个分布是参考分布哪个分布被优化。我已经确认F.kl_div的input参数是 log probabilities。我已经对 logits 做了温度缩放并记录了温度值。我已经用 mask 排除了 padding token。我已经检查参考分布是否存在严格 0 概率并做好了光滑处理。我已经记录了前向 KL 和反向 KL 的基线数值。我准备在验证集上同时看 loss 和生成质量而不是只看 loss。这七项看起来琐碎但每一项都对应一个真实踩坑现场。反向 KL 不是普通 KL 的简单“反向”。它在 OPD 上定义了一种截然不同的约束一个要求覆盖一个允许保守。手撕 Transformer 的真正价值并不只是把网络结构每一层背下来而是在这些看似小到可以忽略的细节里搞清楚你的 loss 到底在优化什么你的模型最终会变成什么样。下次再在代码里看到F.kl_div先别急着复制问自己一句我到底想让哪个分布向哪个分布靠拢如果答案是“学生要靠拢教师”还要继续问是靠拢教师的全部还是只靠拢学生自己选中的那部分。这个问题想明白了反向 KL 才算真正被用对了。