尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

Transformer策略网络在线蒸馏:为什么用反向KL而非正向KL?

Transformer策略网络在线蒸馏:为什么用反向KL而非正向KL? 这次我们不聊一个新模型仓库也不聊一键启动 WebUI而是把一个容易被忽略的数学细节拆开讲清楚Transformer 策略网络在做在线策略蒸馏Online Policy DistillationOPD时为什么很多实现会把目标函数写成反向 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 对多进程采样更友好。Python3.10 或更高。PyTorch2.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_{t1}^{T} \log \pi_\theta(a_t | a_{t}, s) $$因此反向 KL 在轨迹级别可以表示为$$ D_{KL}(\pi_\theta | \pi_t) \mathbb{E}{\tau \sim \pi\theta} \left[ \sum_{t1}^{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, dropout0.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(dim2) 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, dtypetorch.bool, devicex.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, dropout0.1): super().__init__
返回列表