这次我们不聊一个新模型仓库,也不聊一键启动 WebUI,而是把一个容易被忽略的数学细节拆开讲清楚:Transformer 策略网络在做在线策略蒸馏(Online Policy Distillation,OPD)时,为什么很多实现会把目标函数写成反向 KL,而不是习惯里更常见的正向 KL。这个问题在强化学习和模仿学习交叉的地方经常出现,但网上大多数教程都把KL(p||q)和KL(q||p)混着讲,代码里梯度方向稍不注意就反了。
OPD 的场景其实很清晰。先有一个已经收敛的教师策略,可以把它理解为一个用 Transformer 建模动作分布的专家模型;再有一个待训练的学生策略,通常是更小、更快、更适合线上推理的 Transformer。学生进入在线环境,用自己的策略采样轨迹、收集状态;教师不参与交互,只在学生访问到的状态上给出动作分布,学生通过 KL 散度去逼近教师。这里的重点在于:状态分布是学生自己跑出来的,不是从教师分布里抽出来的。所以“期望按谁的分布取”就成了目标函数中最重要的分水岭,这也是反向 KL 在 OPD 里占据主导地位的根本原因。
这篇文章会从公式层面拆解正向 KL 和反向 KL 的差异,推导 OPD 的在线优化目标,然后给出一套最小可复现的 PyTorch 代码,包括一个自注意力策略网络、反向 KL 损失函数、在线蒸馏采样更新循环,以及一套验证反向 KL 是否正常工作的测试方法。本文所有代码都面向训练框架,不涉及对外 API 服务,也不依赖任何第三方模型仓库。
1. 核心能力速览
先给出一张速览表,方便快速判断这篇文章适不适合当前手头的任务。
| 项目 | 说明 |
|---|---|
| 技术类型 | Transformer 策略网络 + 在线策略蒸馏(OPD) |
| 核心数学对象 | 反向 KL 散度 |
| 前置基础 | 概率论、信息论、Transformer 基本结构 |
| 运行框架 | Python 3.10+、PyTorch 2.x、NumPy |
| 硬件要求 | CPU 可跑通最小案例;正式训练推荐 GPU |
| 是否需要现成模型仓库 | 不需要,本文从零实现最小 Transformer 策略 |
| 是否支持 API 服务 | 不涉及,本文提供训练循环 |
| 是否支持批量任务 | 支持,训练和采样都按 batch 组织 |
| 典型输出 | 训练后的学生策略权重 |
| 主要风险 | 模式坍缩、训练不稳定、学生策略熵过早下降 |
从这张表可以直接看出,这篇文章的目标不是给你一个开箱即用的推理工具,而是帮你把“反向 KL 在在线策略蒸馏里到底怎么用”这个问题彻底弄明白,并且能落到代码上验证。如果你只想知道“哪个 API 可以调用”,这篇文章不是你要找的内容。
2. 适用场景与使用边界
反向 KL 不是万能选择。它在 OPD 中的优势来自“在线采样来自学生”这个事实,但一旦场景变成离线数据集蒸馏,或者教师分布和学生分布差异极大,事情就会不同。
适合使用的场景包括:
- 序列决策任务:状态和动作都是 token 序列,Transformer 作为自回归策略网络预测下一个动作 token,适合游戏智能体、对话策略、机器人操作等。
- 在线模仿学习:学生一边和环境交互,一边从教师策略学习,教师只提供动作分布 logits,不提供额外奖励。
- 模型压缩和推理加速:把大教师蒸馏到小学生的过程中,反向 KL 可以更直接地针对学生探索到的状态进行优化。
- 策略迁移:教师和学生的动作空间一致,但模型结构不同,比如教师是 12 层 Transformer,学生是 4 层。
不太适合使用的场景包括:
- 纯离线蒸馏:如果训练数据完全来自静态专家轨迹,动作已经固定下来,此时从教师分布取期望的正向 KL 更容易稳定优化,反向 KL 容易忽略教师分布中的尾部动作。
- 对探索要求极高的稀疏奖励任务:反向 KL 是 mode-seeking 的,学生容易集中在教师高概率区域附近,长期看可能会损失动作多样性。
- 高维连续动作空间:如果动作不是离散 token,而是高维连续向量,反向 KL 需要使用高斯分布或 flow 近似,方差控制会变得复杂。
另外必须强调合规边界。OPD 中教师策略如果是来自某个大型模型的权重,使用前要确认是否有蒸馏、再分发和复制的授权。不要拿未经授权的模型权重去蒸馏线上服务,也不要用于被明确禁止的场景。涉及人脸、声音、版权素材、个人数据时,授权和隐私要求会更高。
3. 环境准备与前置条件
本文不依赖特殊的一键包,只需要一个标准的 Python 深度学习环境。建议按以下清单检查环境。
3.1 系统与版本要求
- 操作系统:Windows、Linux、macOS 均可,Linux 对多进程采样更友好。
- Python:3.10 或更高。
- PyTorch:2.x 版本,支持自动求导和
torch.distributions即可。 - CUDA:如果使用 GPU,推荐 CUDA 11.8 或 12.x;不强制要求。
- 其他依赖:NumPy、tqdm、TensorBoard 或 wandb(用于记录训练曲线)。
安装 PyTorch 时,可以先创建虚拟环境,再按实际机器情况选择安装命令。
python -m venv opd-venv source opd-venv/bin/activate pip install --upgrade pip pip install torch numpy tqdm tensorboard如果是 Windows,激活命令改为:
opd-venv\Scripts\activate如果 GPU 驱动和 CUDA 已就绪,可以在 PyTorch 官网选择对应的 CUDA 版本安装。小规模演示用 CPU 也完全可行,但 Transformer 的序列长度和 batch size 增大后,GPU 差距会非常明显。
3.2 验证环境是否可用
安装完成后,建议先跑一个最小检查脚本,排除环境问题再进入后面的代码。
import torch import torch.nn as nn print(torch.__version__) print("CUDA available:", torch.cuda.is_available())如果输出CUDA available: False,说明当前环境是 CPU 推理。本文的 Demo 规模在 CPU 上可以运行,只是速度会慢一些,不影响理解算法逻辑。
4. 反向 KL 原理与 OPD 推导
这一节是整个文章的核心。如果只想抄代码,可以直接跳到第 5 节;但如果不把KL(p||q)和KL(q||p)的采样方式搞明白,后面调参时会很容易被 NaN 或模式坍缩折磨。
4.1 两个方向的 KL 公式
假设目标分布是 (p),近似分布是 (q)。正向 KL 一般写成:
$$ D_{KL}(p | q) = \mathbb{E}_{x \sim p} \left[ \log \frac{p(x)}{q(x)} \right] $$
反向 KL 写成:
$$ D_{KL}(q | p) = \mathbb{E}_{x \sim q} \left[ \log \frac{q(x)}{p(x)} \right] $$
两者的区别不只是分子分母顺序,而是期望所基于的采样分布完全不同。正向 KL 的样本来自 (p),反向 KL 的样本来自 (q)。在策略蒸馏里,这个区别直接决定了训练数据是如何产生的。
如果用策略网络的角度看,假设学生策略是 (\pi_\theta),教师策略是 (\pi_t),那么反向 KL 可以写成:
$$ D_{KL}(\pi_\theta | \pi_t) = \sum_{a} \pi_\theta(a|s) \log \frac{\pi_\theta(a|s)}{\pi_t(a|s)} $$
这个形式说明,学生只需要在自己的动作分布上计算期望,不需要从教师策略中采样动作。实际上,如果动作空间是离散的,我们甚至可以解析地算出这个 KL,而不需要采样估计。这是反向 KL 在 OPD 中特别好用的原因之一。
4.2 为什么反向 KL 是 mode-seeking
正向 KL 对 (q) 的要求是尽量覆盖 (p) 有概率的所有区域,所以它被称为 mode-covering;反向 KL 则相反,当 (p) 在某个区域概率很低时,只要 (q) 也低,这一项就不会贡献太大损失。反向 KL 的优化结果往往会让学生策略把质量集中到教师概率最高的模式附近,所以它被称为 mode-seeking。
这种特性在在线蒸馏里有直白的解释。学生用当前策略跑出一条轨迹,如果某个动作在学生分布里概率很高,但在教师分布里概率很低,反向 KL 会给出很大的正损失,梯度会立刻把学生从那个方向拉回来。反之,如果学生从来没有采样到某个教师高概率的动作,反向 KL 不会主动去探索它,因为期望是基于 (q) 算出来的。
因此反向 KL 的优势是学习效率高,和在线采样轨迹一致;劣势是探索不足,容易模式坍缩。理解这一点后,后面看训练日志就不会对“学生熵下降太快”感到意外。
4.3 OPD 的目标函数推导
在线策略蒸馏的一个常见目标函数是:
$$ J(\theta) = \mathbb{E}{s \sim d{\pi_\theta}} \left[ D_{KL}(\pi_\theta(\cdot|s) | \pi_t(\cdot|s)) \right] $$
这里的 (d_{\pi_\theta}) 表示学生策略在环境中访问到的状态分布。因为状态来自学生策略,所以外层期望自然按 (d_{\pi_\theta}) 取。如果把内部 KL 展开,就变成:
$$ J(\theta) = \mathbb{E}{s \sim d{\pi_\theta}, a \sim \pi_\theta} \left[ \log \pi_\theta(a|s) - \log \pi_t(a|s) \right] $$
这个结果非常直接。当学生从自己的策略中采样动作 (a) 后,只需要计算学生动作对数概率和教师动作对数概率之差。教师不需要采样完整轨迹,只需要对当前状态给出 logits。
如果换成正向 KL,目标会变成教师策略的期望:
$$ J_{forward}(\theta) = \mathbb{E}{s \sim d{\pi_\theta}, a \sim \pi_t} \left[ \log \pi_t(a|s) - \log \pi_\theta(a|s) \right] $$
问题就在这里:在线场景中,环境交互产生的动作来自 (\pi_\theta),如果损失期望却按 (\pi_t) 取,就需要做重要性采样或者单独用教师策略部署去采集数据,这会显著增加方差和实现复杂度。所以 OPD 通常优先选择反向 KL。
4.4 序列决策中的反向 KL
当策略是自回归 Transformer 时,状态和动作都是 token 序列。对一条轨迹 (\tau = (a_1, a_2, ..., a_T)),对数概率可以展开为:
$$ \log \pi_\theta(\tau|s) = \sum_{t=1}^{T} \log \pi_\theta(a_t | a_{<t}, s) $$
因此反向 KL 在轨迹级别可以表示为:
$$ D_{KL}(\pi_\theta | \pi_t) = \mathbb{E}{\tau \sim \pi\theta} \left[ \sum_{t=1}^{T} \left( \log \pi_\theta(a_t | a_{<t}, s) - \log \pi_t(a_t | a_{<t}, s) \right) \right] $$
这意味着在实现时,不需要计算完整序列分布之间的复杂积分,只需要在每个时间步对比学生和教师对当前步动作的 log-prob。这也是为什么 Transformer 策略和反向 KL 的组合在代码实现上非常自然。
5. Transformer 策略模型与反向 KL 损失实现
从这一节开始进入手撕代码环节。先实现一个最小可用的 Transformer 策略网络,再实现反向 KL 损失函数。
5.1 最小 Transformer 策略网络
为了不让 Demo 过重,我用一个标准 decoder-only 结构,包含因果自注意力、LayerNorm、MLP 和分类头。动作空间用vocab_size表示,block_size表示最大序列长度。
import math import torch import torch.nn as nn import torch.nn.functional as F class CausalSelfAttention(nn.Module): def __init__(self, d_model, n_heads, dropout=0.1): super().__init__() assert d_model % n_heads == 0 self.d_model = d_model self.n_heads = n_heads self.head_dim = d_model // n_heads self.qkv = nn.Linear(d_model, 3 * d_model) self.proj = nn.Linear(d_model, d_model) self.dropout = nn.Dropout(dropout) def forward(self, x): B, T, C = x.shape qkv = self.qkv(x).reshape(B, T, 3, self.n_heads, self.head_dim) q, k, v = qkv.unbind(dim=2) q = q.transpose(1, 2) k = k.transpose(1, 2) v = v.transpose(1, 2) attn = q @ k.transpose(-2, -1) / math.sqrt(self.head_dim) mask = torch.tril(torch.ones(T, T, dtype=torch.bool, device=x.device)) attn = attn.masked_fill(~mask.unsqueeze(0).unsqueeze(0), float("-inf")) attn = F.softmax(attn, dim=-1) attn = self.dropout(attn) out = attn @ v out = out.transpose(1, 2).reshape(B, T, C) return self.proj(out) class TransformerBlock(nn.Module): def __init__(self, d_model, n_heads, dropout=0.1): super().__init__