Transformer策略训练避坑:反向KL散度原理与PyTorch实现详解
不知道你有没有遇到过这种情况:训练一个 Transformer 策略模型时,loss 明明在降,生成内容却越来越单调;或者只是把一行 KL 散度代码的参数顺序换了一下,整个训练曲线就完全变了。我第一次遇到时以为是随机种子的问题,后来把 KL 散度的定义翻出来逐项推,才发现问题的根源往往不在模型结构,而在那个经常被忽略的“方向”。
本文要聊的主题,是 OPD(Online Policy Distillation,在线策略蒸馏)中的反向 KL。如果你所在的项目里 OPD 是其他缩写,不要紧,后面关于正反向 KL 的数学推导、PyTorch 实现和训练避坑内容依然通用。我会从最基础的 KL 散度不对称性讲起,再通过两个可运行的 PyTorch 实验,带你把反向 KL 在 Transformer 策略蒸馏中的行为彻底“手撕”清楚。
1. 背景与核心概念
1.1 Transformer 为什么会被当作策略模型
传统的策略模型可以是 CNN、RNN、MLP,但在序列决策和生成任务里,Transformer 凭借自注意力机制,能够建模长距离依赖,并且在处理离散 token 序列时非常自然。比如在 RLHF、决策 Transformer、在线策略蒸馏这类场景中,我们经常看到一个 Decoder-only 的 Transformer 直接输出下一个 token 的概率分布,这个分布就是策略。
当模型规模和训练数据足够大时,Transformer 策略往往比小网络更有表达能力,但训练也更容易不稳定。模型输出的是一个高维离散分布,分布之间如何比较、如何约束,就成了训练稳定性的关键。
1.2 OPD 到底在做什么
OPD 最常见的含义是 Online Policy Distillation,也就是在线策略蒸馏。简单说,我们有一个教师策略 P 和一个学生策略 Q,两者通常都是 Transformer。教师可能是训练了很久的模型,也可能是某个效果更好的上一个版本。学生需要不断学习教师的知识,让自己输出分布尽量接近教师。
“在线”体现在训练过程中,学生模型的采样轨迹会不断变化。这意味着学生今天采样的分布和昨天采样的分布可能差别很大,蒸馏目标如果设计不好,训练很容易振荡甚至崩溃。这也是为什么我们不能随手拿一个交叉熵就开训,必须仔细选择分布之间的距离度量。
1.3 KL 散度为什么是核心
KL 散度(Kullback-Leibler Divergence)用于度量两个概率分布之间的差异。教师分布 P 和学生分布 Q 之间的 KL 散度越大,说明两者差异越大。策略蒸馏本质上就是最小化这个差异。
但这里有一个非常容易踩坑的点:KL 散度不是对称的。
更准确地说,它们大部分情况下都不相等。一个是“以 P 为基准”,另一个是“以 Q 为基准”。这两个方向在优化时对应完全不同的行为,这正是“反向 KL”这个说法的来源。要知道为什么 OPD 中常用反向 KL,需要先把这两个方向的数学性质搞清楚。
2. 从 KL 散度的不对称性说起
2.1 正反向 KL 的数学定义
假设两个离散概率分布 P 和 Q,取值空间为 x。
正向 KL 一般写作:
反向 KL 写作:
如果只看公式,两者只是 P 和 Q 的位置交换。但在优化一个分布时,这个位置交换会带来完全不同的梯度行为。为了叙述方便,下文中我们默认要优化的是学生分布 Q,教师分布 P 固定不变。
2.2 正向 KL 是 mean-seeking
先看 KL(P || Q)。它的每一项都乘以 P(x)。也就是说,只有 P(x) 比较大的位置才会对 loss 产生显著影响。如果某个位置 P(x) 很大而 Q(x) 很小时,log(P(x)/Q(x)) 会非常大,梯度会强烈推动 Q 在该位置增加概率。
因此,优化 KL(P || Q) 时,Q 会尽量覆盖 P 的所有高概率区域。哪怕 P 是双峰分布,Q 也会尝试把两个峰都覆盖住。因为如果不覆盖,loss 就压不下去。
这种特性被称为:
- mean-seeking
- zero-avoiding
- 覆盖模式
意思是 Q 倾向于铺开,试图用自身分布包住 P 的主要概率质量。
2.3 反向 KL 是 mode-seeking
再看 KL(Q || P)。它每一项乘以 Q(x)。也就是说,Q 在哪个位置分配概率,哪个位置就会进入 loss。如果 Q 在某个位置 x 分配了概率而 P(x) 非常小甚至接近 0,那么 log(Q(x)/P(x)) 会非常大,导致 loss 爆炸。
所以优化反向 KL 时,Q 会倾向于只在 P 概率较高的位置附近分配概率,绝不轻易探索 P 没有覆盖的区域。
这种特性被称为:
- mode-seeking
- zero-forcing
- 模式锁定
当 P 是双峰分布而 Q 是单峰分布时,Q 最终会锁在其中一个峰附近,而不是同时覆盖两个峰。反向 KL 的优化结果具有一种“锐化”倾向,会让分布变得更集中。
2.4 两种 KL 行为对比
| 对比维度 | 正向 KL(P || Q) | 反向 KL(Q || P) |
|---|---|---|
| 优化 Q 时乘以什么权重 | P(x) | Q(x) |
| 对 Q 在 P 低概率区域的表现 | 惩罚较小 | 惩罚极大 |
| 最终分布倾向 | 覆盖所有模式 | 锁定某个模式 |
| 常被称为 | mean-seeking | mode-seeking |
| 对分布熵的影响 | 让学生熵变大或保持覆盖 | 让学生熵变小、更锐化 |
这里要特别提醒:上面说的是“优化 Q 时”的行为。如果优化的是 P,那么两个方向的性质会互换。很多实现出错,就是因为在写法上没有搞清楚当前到底在优化哪个分布。