大模型训练里有一个经常被忽略的区分:技能到底应该存在哪一层。是把推理步骤写进 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-Policy:DQN 的经验回放
DQN 是 off-policy 的典型代表,因为 Q-learning 更新的是状态动作价值函数,训练时使用的行为策略可以带探索噪声,而目标值用的是贪心策略。只要存的 transition 能覆盖足够的 (s, a, r, s'),反复回放也能收敛。
| 对比维度 | PPO | DQN |
|---|---|---|
| 数据来源 | 必须由当前策略采样 | 可以由任意行为策略采样 |
| 数据复用 | 每轮迭代后旧数据作废 | 可以放入经验回放反复使用 |
| 重要性修正 | 通过 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_epochs=4, clip_epsilon=0.2, gamma=0.99, lam=0.95, skill_loss_weight=0.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_policy:PPO clip loss,负责优化 token 层面的策略。
- L_value:价值网络拟合误差,负责给 GAE 提供基线。
- L_skill:技能分类 loss,负责让模型学会"当前状态应该使用什么抽象技能"。
三者的权重配比需要实验调整。一个常见的问题是 skill loss 权重过高会挤压策略 loss,导致模型只会分类、不擅长生成;权重过低则技能信息没有被有效压进权重。可以参考的经验值是skill_loss_weight从 0.05 到 0.2 之间搜索,具体以验证集行为为准。
5. 环境准备与最小实现
5.1 环境与依赖
如果原始材料没有给出明确版本,落地前要先确认自己环境的 CUDA 和显卡是否匹配。下面是一个常见的组合,用于说明思路:
| 软件 | 版本建议 | 用途 |
|---|---|---|
| Python | 3.10+ | 运行环境 |
| PyTorch | 2.1+ | 张量计算和自动求导 |
| Transformers | 4.38+ | 模型加载和 tokenizer |
| TRL | 0.9+ | 可选的 PPO Trainer 封装 |
| CUDA | 11.8 或 12.1 | GPU 加速 |
学习阶段建议先用 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,就能把"技能蒸馏"从论文思路变成一条可以持续迭代的训练管线。