行业资讯
📅 2026/8/29 13:20:23
技能蒸馏进权重:On-Policy自蒸馏与抽象特权信号
大模型训练里有一个经常被忽略的区分技能到底应该存在哪一层。是把推理步骤写进 prompt让模型每次生成都照着执行还是通过训练把技能压进权重让模型在不给提示的情况下也天然具备这种能力。标题中的研究思路选择的是后者——把技能蒸馏进权重而不是提示词并围绕 on-policy 自蒸馏和抽象技能特权信号来设计训练流程。这篇文章会把这个思路拆成四个关键概念用 PPO 的 on-policy 机制解释为什么蒸馏过程需要重新采样再给出一套可落地的训练框架、最小实现和排查路径。读完以后你能理解技能蒸馏进权重和用 prompt 堆能力之间的本质差异也能在小型模型上先跑通一个 on-policy 自蒸馏循环。1. 先想清楚一个问题技能存在提示词里还是存在权重里1.1 技能写在 prompt 里问题出在哪在 LLM 应用里最直接的给模型注入技能方式是把技能写进 prompt给几个 few-shot 示例让模型模仿在 system prompt 里写请分步骤推理或者直接拼接一段 CoT 示例引导中间思考。这种方式很方便不需要训练改一行文本就能换一套行为适合快速验证。但它在工程上有几个绕不过去的代价每次调用都要重复输入技能描述推理 token 增多成本和时延上升。技能只是临时指令模型可能时灵时不灵尤其当上下文变长或指令被其他内容覆盖时。换一个基座模型prompt 需要重新调做模型量化、蒸馏到小模型时prompt 并不能保证被保留。最重要的一点模型没有真正学会这个技能它只是被提示词临时引导。换句话说技能写在 prompt 里本质是运行时借来的能力权重里没有留下痕迹。1.2 把技能压进权重意味着什么蒸馏进权重是指通过梯度更新让模型参数的分布形态发生改变从而在没有任何额外提示的情况下也能以较高概率输出带有技能特征的中间结果。比如一个模型经过抽象技能蒸馏后遇到数学题会先写约束条件再计算遇到代码任务会先列测试再写实现。这些行为不是靠 prompt 指令临时控制的而是由参数内化出来的。这种方式适合对稳定性和推理开销敏感的生产场景。缺点也很明显需要训练数据、需要评估流程、每次调整技能都要重新训练或微调迭代速度远慢于改 prompt。下面用一张表对比两条路线的差异对比维度技能放在 Prompt 里技能蒸馏进权重是否需要训练不需要需要推理 token 开销有额外开销基本无行为稳定性依赖模型遵循指令更稳定更内化迁移到其他基座需要重新调 prompt随权重迁移迭代速度快改文本即可慢需要重新训练可解释性能看到提示内容只能通过行为验证理解这个对比后才能理解标题为什么要强调distill into weights, not prompts。它不是否定 prompt 的价值而是指出一个更持久的技术目标。1.3 从知识蒸馏到自蒸馏老师不再是外部模型传统知识蒸馏Knowledge Distillation是典型的 teacher-student 结构大模型输出软标签或中间层特征小模型去拟合。这个过程中老师通常是固定的外部模型学生学习的是老师看到的数据分布。自蒸馏Self-Distillation的关键区别在于老师信号和学生来自同一个模型家族甚至直接来自学生自身。常见形态有两种结构自蒸馏用模型更深层的输出监督浅层或者用 EMA 维护的慢版本模型作为老师。数据自蒸馏让模型自己生成数据再用某种信号标注器给这些数据打上额外信息最后模型在自己的生成结果上继续训练。标题中的方案本质是第二种并且强调了一个重要约束生成数据和更新权重必须保持在同一个策略分布上也就是 on-policy。2. 拆解标题里的四个关键词这一节把标题里最重要的四个术语逐个拆开因为它们各自都对应一组工程选择。2.1 Self-Distillation谁的输出在教谁自蒸馏里最容易搞混的是教的方向。在数据自蒸馏中学生先生成一批 rollout然后一个信号生成器可能是同一个模型、一个冻结的大模型、或者一个规则评分器对 rollout 进行标注最后学生用自己的 rollout 和自己的标注做梯度更新。这个流程里老师不是直接给出一段标准答案让学生背而是给学生生成的样本打分、归类、补充过程性信息。学生学到的是我的输出哪些地方值得加强而不是换一种完全不属于我的输出风格。自蒸馏形态老师信号来源典型实现方式结构自蒸馏网络深层 / EMA 模型深层 logits 监督浅层输出数据自蒸馏模型自身生成 标注器rollout 采样 信号监督 策略更新2.2 On-Policy样本必须来自当前策略on-policy 的中文直译是在策略内含义是更新策略所用的数据必须由当前版本的策略采样得到。这句话初看像废话但在强化学习和蒸馏场景里非常重要。策略参数每更新一次模型输出分布就会变化。上一版权重采样的样本对当前权重来说已经不是当前分布下的样本了分布已经发生了偏移。如果蒸馏过程使用的是固定老师生成的数据或者使用学生上一个 checkpoint 生成的旧数据那就是 off-policy 式的学习。它不是不能用但会出现明显的分布不匹配问题模型学会的是模仿旧版本自己或模仿老师而不是优化当前自己的行为。标题强调 on-policy本质上是想让梯度更新和采样分布严格对齐。2.3 Privileged Signals特权信号是什么为什么需要特权信号privileged signals指那些训练阶段可以使用、但推理阶段不一定存在的额外信息。这个词在机器人学习里很常见比如训练时可以使用真实物体位置、地面真实状态推理时模型只能依赖自己的传感器输入。迁移到大模型场景中特权信号可以有很多具体形态生成过程中每个步骤的过程奖励分数。一个冻结大模型给学生 rollout 标注的抽象技能分类 ID。更优的下一步中间思路或局部动作提示。只能在训练集里看到的真实标签或目标答案。这些信号的作用是引导梯度方向。学生不需要在推理时输出这些信号也不需要从输入中推断它们它们只负责把模型推向正确的行为区域。标题把它和abstract skills组合在一起表达的是用抽象技能作为特权信号而不是用简单的好/坏标量作为唯一反馈。2.4 Abstract Skills抽象技能和 token 级标注的区别一个很常见的错误做法是让老师模型直接生成一段更好的续写然后让学生做监督学习拟合。这是 token 级的知识蒸馏问题在于 token 级目标太细、太容易过拟合表面形式学生学到的是表面句式而不是问题的结构化处理方式。抽象技能位于比 token 更高的层次。它不是下一句话说什么而是这类问题应该采取什么处理策略。举例来说数学题先提取约束条件再列出需要满足的不等式最后计算。代码任务先写最简测试再实现函数最后运行测试修正。问答任务先判断问题类型再决定是否需要检索外部知识。在实现上抽象技能可以被编码成离散技能 ID 或连续技能向量。学生模型需要在每个状态下预测自己正在使用哪个技能同时根据该技能规划后续动作。这个状态到技能的映射就是最终沉淀进权重的核心能力。3. 理解 On-Policy从 PPO 为什么是 on-policy 说起热搜里有一个高频问题为什么说 PPO 是 on-policy。这个问题和自蒸馏流程直接相关本节单独展开。3.1 一句直观解释PPO 是 on-policy 算法的原因只有一个它的目标函数里用了一个重要性比值而这个比值只有在数据确实由旧策略采样时才有统计意义。PPO 的优化目标可以简化成L min( ratio * A, clip(ratio, 1-ε, 1ε) * A )其中ratio π_new(a|s) / π_old(a|s)如果一条样本不是由 π_old 采样的那这个比值本身就是错的整个 clip 逻辑也失去意义。3.2 PPO 的数据流PPO 的标准流程是策略 π_old 与环境交互采样一批轨迹。计算每条轨迹的优势估计通常用 GAE。在固定这一批数据上做多轮梯度更新但每一轮都会检查新旧策略的 KL 散度并用 clip 限制更新幅度。更新完成后旧数据作废必须重新采样。所以on-policy不仅是理论属性还是工程约束。它决定了你不能像 DQN 那样把大量旧样本塞进 replay buffer 反复利用。3.3 对比 Off-PolicyDQN 的经验回放DQN 是 off-policy 的典型代表因为 Q-learning 更新的是状态动作价值函数训练时使用的行为策略可以带探索噪声而目标值用的是贪心策略。只要存的 transition 能覆盖足够的 (s, a, r, s)反复回放也能收敛。对比维度PPODQN数据来源必须由当前策略采样可以由任意行为策略采样数据复用每轮迭代后旧数据作废可以放入经验回放反复使用重要性修正通过 ratio 和 clip 修正不需要专门修正收敛稳定性依赖采样质量和 KL 约束依赖回放缓冲和 target 网络理解这个差异后回到蒸馏场景如果在自蒸馏中直接复用旧 checkpoint 生成的 rollout 来训练新 checkpoint实际上就是在做 off-policy 学习需要额外补偿分布偏移如果严格重新采样则是 on-policy更贴近 PPO 的设计哲学。3.4 On-Policy 蒸馏和 PPO 的关系标题中的 on-policy self-distillation在实现上通常直接套用 PPO 框架。学生模型就是策略网络它生成 rollout然后一个特权信号模块给每个 token 或每个 step 打信号再计算 GAE 优势最后用 PPO loss 更新。整个过程和 RLHF 里的 PPO 训练非常相似区别在于奖励信号不只来自结果还来自抽象技能标注和过程信号。额外的 skill loss 会把状态到技能的映射也压进权重。所以理解 PPO 为什么是 on-policy是理解这套自蒸馏方法的前提。4. 方法主框架抽象技能作为特权信号的训练流程4.1 训练循环总览整个训练循环可以拆成五个阶段当前学生策略在任务 prompt 上采样 rollout。信号生成模块对 rollout 做标注输出三类信号技能 ID、过程奖励、特权提示。计算 GAE 优势。在旧 rollout 上做多轮 PPO 更新同时用 skill loss 训练技能预测头。更新完成后丢弃旧样本回到第 1 步。这个循环最核心的设计是信号永远基于学生当前生成的 rollout 计算而不是单独生成一套标准答案。# 一次 on-policy 自蒸馏迭代的伪代码 def single_iteration(actor, skill_head, signal_fn, tokenizer, prompts, optimizer, ppo_epochs4, clip_epsilon0.2, gamma0.99, lam0.95, skill_loss_weight0.1): # 1. 当前策略采样 rollout rollouts collect_rollouts(actor, tokenizer, prompts) # 2. 信号模块生成特权信号 rollouts signal_fn.annotate(rollouts) # 3. 计算 GAE 优势 advantages compute_gae(rollouts[rewards], rollouts[values], gamma, lam) # 4. 在旧样本上做多轮 PPO 更新 for _ in range(ppo_epochs): policy_loss clip_loss(rollouts, actor, clip_epsilon, advantages) value_loss mse_loss( actor.compute_values(rollouts[states]), rollouts[returns] ) skill_loss cross_entropy( skill_head(rollouts[states]), rollouts[skill_ids] ) total_loss (policy_loss 0.5 * value_loss skill_loss_weight * skill_loss) optimizer.zero_grad() total_loss.backward() optimizer.step() # 5. 旧样本作废下一轮重新采样代码块后的关键点ppo_epochs不能设置太大因为 on-policy 属性决定了这批数据只能临时使用复用过猛会让策略分布和采样分布偏差变大。4.2 信号生成模块信号生成模块是整套框架的老师它的输入是学生生成的完整轨迹输出是结构化的特权信号。根据实现成本可以分成几个层级信号类型内容监督形式实现成本技能 ID抽象技能类别离散分类 CE loss低需要预定义技能库技能 Embedding连续技能向量回归或对比学习中过程奖励每个 step 的质量分价值函数拟合中高需要过程标注特权提示更优的中间思路生成式 loss高需要强教师模型实际项目里建议从技能 ID 加过程奖励起步。先有一个稳定的技能分类体系再逐步加入连续向量和特权提示。4.3 目标函数组合训练总损失由三部分组成L_policyPPO clip loss负责优化 token 层面的策略。L_value价值网络拟合误差负责给 GAE 提供基线。L_skill技能分类 loss负责让模型学会当前状态应该使用什么抽象技能。三者的权重配比需要实验调整。一个常见的问题是 skill loss 权重过高会挤压策略 loss导致模型只会分类、不擅长生成权重过低则技能信息没有被有效压进权重。可以参考的经验值是skill_loss_weight从 0.05 到 0.2 之间搜索具体以验证集行为为准。5. 环境准备与最小实现5.1 环境与依赖如果原始材料没有给出明确版本落地前要先确认自己环境的 CUDA 和显卡是否匹配。下面是一个常见的组合用于说明思路软件版本建议用途Python3.10运行环境PyTorch2.1张量计算和自动求导Transformers4.38模型加载和 tokenizerTRL0.9可选的 PPO Trainer 封装CUDA11.8 或 12.1GPU 加速学习阶段建议先用 1B 参数以下的模型跑通循环比如小型 LLaMA、Qwen 或 GPT-2。显卡用单张 24GB 显存即可。生产环境再考虑更大的模型和多卡并行。5.2 最小代码结构一个可复现的最小项目可以按下面结构组织skill_distill/ ├── configs/ │ └── ppo.yaml ├── data/ │ └── tasks.py ├── models/ │ ├── policy.py │ └── skill_head.py ├── teachers/ │ └── skill_annotator.py ├── algos/ │ ├── ppo_buffer.py │ └── trainer.py └── scripts/ └── run_train.py# configs/ppo.yaml model_name: Qwen/Qwen2-0.5B-Instruct batch_size: 8 max_steps: 64 ppo_epochs: 4 clip_epsilon: 0.2 gamma: 0.99 lam: 0.95 skill_loss_weight: 0.1 learning_rate: 1e-65.3 核心代码片段先写 rollout 采集。这里的关键是记录每个 token 的 log_prob 和状态值供后续 PPO 更新使用。# algos/ppo_buffer.py 片段记录采样信息 def append_step(self, state, action, logp, value, reward): self.states.append(state) self.actions.append(action) self.logp.append(logp) self.values.append(value) self.rewards.append(reward)然后是 PPO 的 clip loss。注意这里使用的 logp_old 来自采样时记录的值logp_new 来自当前权重重新前向计算。# algos/trainer.py 片段PPO clip loss import torch def clip_loss(rollouts, actor, clip_epsilon, advantages): states rollouts[states] actions rollouts[actions] logp_old rollouts[logp] logp_new actor.log_prob(states, actions) ratio (logp_new - logp_old).exp() surr1 ratio * advantages surr2 ratio.clamp(1.0 - clip_epsilon, 1.0 clip_epsilon) * advantages return -torch.min(surr1, surr2).mean()最后是技能预测头。技能头接收状态表示输出技能 ID 的概率分布和策略共享大部分底层参数只保留一个独立分类头。# models/skill_head.py 片段 import torch.nn as nn class SkillHead(nn.Module): def __init__(self, hidden_size, num_skills): super().__init__() self.classifier nn.Linear(hidden_size, num_skills) def forward(self, hidden_states): return self.classifier(hidden_states)5.4 运行与验证验证分三步走第一步确认训练循环能跑通不要求效果好只看 loss 是否正常下降、显存是否溢出、采样和更新是否交替执行。第二步固定一个小任务集比如 100 道需要多步推理的题观察三个指标策略平均熵是否保持在一个合理区间、KL 散度是否没有突然暴涨、技能分类准确率是否在上升。第三步做一次去 prompt 评测。训练结束后把测试 prompt 里的技能描述和 few-shot 例子全部移除只保留问题本身看模型是否还能产出带技能特征的中间步骤。预期结果应该是未蒸馏模型在去掉技能 prompt 后效果明显下降而完成蒸馏的模型下降幅度显著更小。这就是技能进入权重的直接证据。6. 常见问题与排查路径6.1 用表格快速定位问题下面这张表整理了训练中最高频出现的四类问题问题现象常见原因检查方式处理建议loss 不降或震荡数据分布和信号不匹配查看 rollout 分布、KL 散度确保每轮更新后重新采样减少 ppo_epochs技能预测退化成常数技能粒度太粗或类别不平衡查看 skill 分布熵调整技能类别、增加平衡采样、加大熵正则去掉 prompt 后效果变差模型仍依赖 prompt 触发对比有 / 无 prompt 的评测结果训练时随机删除技能描述前缀奖励被刷高但行为变差reward hacking分析高奖励样本的具体行为改用过程奖励加入技能覆盖度指标6.2 典型坑详解坑一把旧 rollout 直接塞进下一轮训练。这样做的结果是分布偏移学生更新的方向和新采样分布不一致训练会越来越不稳定。正确做法是每一轮 PPO 更新结束后清空 buffer下一轮用新权重重新采样。坑二特权信号泄露到推理路径。如果训练时把教师给出的更好的下一步思路拼进了输入模型的状态里推理时却没有这个信号模型会依赖一个不存在的输入表现会断崖式下降。要保证特权信号只参与 loss 计算不进入 token 输入序列。坑三技能分类头学会偷懒。当技能库类别太少时模型不需要区分状态就能猜对当信号标注噪声太大时技能头会趋向于输出分布均匀或退化成常数。解决方式是设计层级化技能库并在 skill loss 上加熵正则惩罚过强的确定性输出。坑四直接在大模型上调参。如果一开始就在 7B 或 70B 模型上跑单次迭代成本和调试周期会非常高。建议先在小模型上把整条链路调通确认信号设计合理后再放大。6.3 排查链路按下面顺序排查能覆盖绝大多数问题输入是否正确任务 prompt、tokenizer 特殊 token、padding 是否正确。采样分布是否正常观察生成文本的长度、重复率、熵值。信号是否合理随机抽 20 条 rollout人工检查技能 ID 和过程奖励是否有明显错误。优势计算是否正确检查 GAE 返回值是否出现极端数值。loss 各分量是否在合理量级如果 policy loss 比其他两个大很多需要调整权重。评测是否客观去掉 prompt 后对比不能只看训练集指标。7. 工程实践建议与扩展方向7.1 学习环境与生产环境的差异学习环境求跑通生产环境求可控。两套环境的差异要区分开关注点学习环境生产环境模型规模1B 以下按业务需求选择显卡单卡 24GB多卡 / 推理集群日志print wandb结构化日志、指标监控回滚不需要需要 checkpoint 版本管理评估人工抽样自动化评测集 线上灰度数据小样本任务全量数据、数据版本记录7.2 训练前检查清单每次启动训练前建议先过一遍下面的清单技能库是否有明确类别定义类别之间是否相互独立。特权信号是否严格只出现在 loss 中不出现在输入序列中。采样 buffer 是否在每轮更新后被清空。PPO 的 clip_epsilon 和 ppo_epochs 是否已按小模型调试过。是否记录了策略熵、KL 散度、技能分类准确率、技能覆盖度四个指标。评测集是否包含无 prompt 条件用于判断技能是否真正进入权重。是否预留了回滚 checkpoint。7.3 扩展方向这套框架在真实项目里可以往几个方向扩展。第一个方向是技能库的构建。从少量手写技能起步逐步用聚类方法从高质量数据里自动发现技能类别并把离散技能 ID 升级为连续技能向量这样技能之间的相似性也能被利用。第二个方向是持续学习。把不同任务逐步蒸馏进同一个模型时抽象技能可以作为任务间的共享组件降低灾难性遗忘的影响。第三个方向是和 prompt 的混合使用。权重内化技能不等于完全放弃 prompt生产里常见的做法是核心技能要求模型无条件执行所以压进权重临时性偏好或业务策略用 prompt 控制。两者结合既保证稳定性又保留灵活性。最后一个思路是把 on-policy 自蒸馏复用在你已有的 PPO 训练框架上。如果团队已经跑通 RLHF只需要在 PPO 循环里额外加一个 skill head 和一个 skill loss就能把技能蒸馏从论文思路变成一条可以持续迭代的训练管线。