知识蒸馏进入大模型时代后,一个研究方法论层面的问题被摆到了台面上:当教师模型不再作为离线标签生成器,而是顺着学生模型自己采样的分布,逐 token 给出概率目标,这种 on-policy 蒸馏到底是在“蒸馏”,还是在做另一种策略强化更新?围绕这个提问,相关论文会重新拆解 on-policy 蒸馏的优化目标,并给出不依赖外部教师监督的 OPSA(On-Policy Self-Alignment,在策略自对齐)方案。这篇文章沿着这条研究主线展开,先说明离线蒸馏和在线蒸馏的本质差异,再用损失函数拆解为什么 on-policy 蒸馏可能会被误读,最后给出 OPSA 的思路、最小实现和验证方案。
这篇整理适合正在做大模型偏好对齐、推理能力蒸馏、RLHF/DPO 改造,或者想复现“on-policy 蒸馏是否真的在蒸馏”这类实验的读者。读完你会得到一套判断框架:什么时候蒸馏真的是在迁移教师知识,什么时候只是在借用教师反馈做策略优化,以及一个无需监督标签的自对齐训练循环如何搭建。
1. 先界定“蒸馏”:固定数据学习教师,和在策略采样学习教师不是一回事
1.1 经典蒸馏解决什么问题
知识蒸馏最初要解决的是模型压缩问题。一个容量很大的教师模型已经在一个任务上学到了输入到输出的映射,学生模型希望用更少的参数逼近这个映射,但直接拿硬标签训练可能丢失教师输出的不确定性信息。于是蒸馏使用教师的软化概率作为训练目标:
[ \mathcal{L}{\mathrm{KD}} = \mathbb{E}{x \sim \mathcal{D}}\left[ \mathrm{KL}\left(q_T(\cdot|x) | \pi_\theta(\cdot|x)\right) \right] ]
其中 (q_T) 是教师模型在输入 (x) 上的输出分布,(\pi_\theta) 是学生模型。温度系数加高后,教师分布中小概率但存在语义关联的类别也能给学生提供学习信号。这个过程的两个关键点是:
- 数据分布 (\mathcal{D}) 固定不变,通常来自一个预先收集好的训练集;
- 教师只负责产生目标,不会因为学生变化而重新生成数据。
在这种设定下,训练目标和模仿学习关系非常明确:学生在固定输入上调整自己的条件分布,让它逼近教师的条件分布。
1.2 on-policy 蒸馏改变了哪些东西
强化学习和语言模型对齐里说的 on-policy,核心是“采样分布来自当前策略本身”。学生的推理样本由学生自己生成,状态分布、动作分布、整条轨迹分布都随学生参数一起演变。教师模型不再是一个静态标签提供器,而是在学生访问到的每个中间状态上给出评价或目标。
于是 on-policy 蒸馏的损失可以写成:
[ \mathcal{L}{\mathrm{KD}}^{\mathrm{on}} = - \mathbb{E}{x \sim \mathcal{D}} \left[ \mathbb{E}{y \sim \pi\theta(\cdot|x)} \left[ \log q_T(y|x) \right] \right] ]
直观上看,表达式里仍然有教师分布 (q_T),所以很多实现会把这段代码放进“蒸馏”目录下。但注意期望的采样来源已经变成学生策略 (\pi_\theta),这和传统蒸馏中“学生用固定数据学习教师”并不相同。
1.3 三种常见“蒸馏式训练”的差别
| 训练方式 | 样本来源 | 目标来源 | 数据分布随学生变化 | 更像什么 |
|---|---|---|---|---|
| 离线蒸馏 | 固定训练集 | 教师预计算 soft label | 否 | 监督学习/模仿学习 |
| on-policy 蒸馏/教师打分 | 学生实时生成 | 教师对生成样本打分 | 是 | 奖励塑形后的策略优化 |
| 自蒸馏/自对齐 | 学生实时生成 | 学生自己的一致性信号 | 是 | 无监督策略自改进 |
这张表说明了论文标题里那个看上去像反问的标题为什么值得较真:如果一个流程的样本来源、数据分布、损失本质都已经和经典蒸馏不同,那么继续把它称作“蒸馏”,会掩盖它实际在做的事情。尤其当论文要论证“学生提升来自教师知识迁移”时,on-policy 的采样机制本身就会引入一个更强的混淆变量:策略优化。
2. 从损失函数看:on-policy 蒸馏在梯度上更接近策略梯度
2.1 用一个单步决策例子拆解
为了把机制讲清楚,先不考虑多步推理的轨迹长度,只看单步动作分布。设学生策略是 (\pi_\theta(a|s)),教师条件分布是 (q_T(a|s))。on-policy 蒸馏的目标可以理解成最大化学生采样动作被教师认可的程度:
[ \mathcal{J}(\theta) = \mathbb{E}{s \sim \mathcal{D}, a \sim \pi\theta(\cdot|s)} \left[ \log q_T(a|s) \right] ]
对这个目标求梯度,使用 score function 技巧后:
[ \nabla_\theta \mathcal{J}(\theta)
\mathbb{E}{s, a \sim \pi\theta} \left[ \left( \log q_T(a|s) - b(s) \right) \nabla_\theta \log \pi_\theta(a|s) \right] ]
这里的 (b(s)) 可以是任意只依赖状态不依赖动作的基线,因为 (\mathbb{E}{a \sim \pi\theta}[\nabla_\theta \log \pi_\theta(a|s)] = 0),加上基线不会改变期望梯度。
这个形式和强化学习里的策略梯度几乎一致:(\log q_T(a|s)) 扮演奖励函数,学生自己采到的动作扮演 rollout,(\nabla_\theta \log \pi_\theta(a|s)) 是提升该动作概率的方向。教师概率高的动作会被推高,教师概率低的动作会被压低。
2.2 为什么不能说它“一定没在蒸馏”
需要说明一点:on-policy 蒸馏和奖励塑形并不完全等价于随机乱学。当学生采样到的动作恰好覆盖了教师高概率区域时,最大化教师对数概率确实会把学生往教师方向推。在任务分布比较窄、动作空间比较小、教师评估又很稳定的场景里,on-policy 蒸馏能起到类似对齐的效果,也能提升下游分数。
真正的问题在于效果的归属。假设一个学生模型经过 on-policy 蒸馏后分数提升,我们无法判断提升来自:
- 学生真的学会了教师的推理规则和决策偏好;
- 还是仅仅因为学生把自身访问分布内“教师觉得好”的动作概率调高了;
- 还是因为采样和优化本身带来了额外的探索收益。
离线蒸馏没有这个问题,因为学生始终在固定输入上做分布匹配,教师知识是学生改变行为的直接目标。on-policy 蒸馏由于状态访问分布也被优化,学生可以走捷径:它可能只提升自己在高教师置信状态上的概率,而完全不需要理解教师在其他状态下的行为。
2.3 实验结果应该怎么设计才不能自证
如果论文想回答“on-policy 蒸馏是否真的在蒸馏”,最简单的做法是加三组对照:
| 实验变体 | 采样来源 | 目标 | 理论实质 |
|---|---|---|---|
| Offline-KD | 固定教师生成数据 | 教师 soft label | 标准知识蒸馏 |
| On-policy-KD | 学生 rollout | 教师对数概率 | 教师奖励形式的策略优化 |
| Teacher-Reward-RL | 学生 rollout | 一个奖励函数等于教师概率 | 常规策略优化 |
| OPSA | 学生 rollout | 无外部教师的一致性信号 | 无监督自对齐 |
如果 On-policy-KD 的指标与 Teacher-Reward-RL 非常接近,却和 Offline-KD 在状态覆盖、行为距离、分布外泛化上差异明显,那就说明 on-policy 变体本质更接近策略优化,而不是知识迁移。这种对照设计是复现“蒸馏是否真的发生”的关键。
2.4 一个容易被误解的点
很多人会把 on-policy 蒸馏等价于 DAgger。DAgger 虽然是 on-policy 采样专家轨迹,但它的核心是让专家在学生的状态分布上标注动作,然后监督学习这些“状态-动作对”。DAgger 的目标是匹配专家条件策略,只是抽样分布换成了学生。而很多语言模型场景里的 on-policy 蒸馏并不收集“教师动作”,只把教师概率当成 soft reward 反馈给学生,这两种机制需要严格区分。
判断方法很简单:看学生训练时,教师是否在学生生成的每个位置上给出了一个完整的目标分布,并且损失里是否包含“把学生分布拉到该目标分布”的 KL 项。如果教师只是给教师自己预测的 token 打分,学生学到的其实是奖励最大化。
3. OPSA 的设计:不使用外部监督,如何还用 on-policy 改进模型
3.1 OPSA 的动机来自哪里
顺着上一节的判断,on-policy 蒸馏的一个尴尬之处在于:如果教师信号本质是奖励,那么它既没有充分利用教师的条件分布知识,又保留了策略优化带来的方差和数据分布偏移问题。那有没有可能跳过教师,直接利用模型自身在同一个提示下产生的多条采样,构造一个无需外部监督的对齐信号?这正是 OPSA 想解决的问题。
OPSA 里的 O 是 On-Policy,P 是 Policy,S 是 Self,A 是 Alignment。它不在每一轮从教师模型读取概率,而是让当前模型产生多条候选输出,再用候选输出之间的自洽性决定哪些输出更值得被强化。整个过程不需要人类标注、不需要奖励模型、也不需要教师 soft label。
3.2 自洽分数如何构造
假设当前模型对提示 (x) 采样了 (K) 条回答 (y_1, y_2, \dots, y_K)。对于有确定答案的任务,可以先做答案抽取,再用答案字符串匹配或语义相似度判断两条回答是否一致。
对第 (k) 条回答,可以定义它的一致性权重:
[ c_k = \frac{1}{K-1} \sum_{j \neq k} \mathbb{1}\left[ \mathrm{answer}(y_k) = \mathrm{answer}(y_j) \right] ]
这个值表示该回答和其他采样一致的比例。如果多数模型自己生成的结果都收敛到同一个答案,那么属于该答案簇的样本一致性高,模型更愿意保留这种推理模式。
在实现时,常用两种变体:
- 硬簇权重:只给答案出现次数最多的样本权重 1,其余为 0;
- 软簇权重:按簇大小归一化,让样本权重等于它所在簇的占比。
软簇权重的梯度更平滑,不容易因为单次采样噪声导致某一条回答被暴力抬高。实际项目里建议先用软权重,等损失稳定后再判断是否需要切换。
3.3 OPSA 的目标函数
OPSA 的每个更新步骤等价于做一次 KL 正则的 on-policy 策略优化。首先用当前策略冻结采样,得到一批样本;然后计算每个样本的 self-consistency reward;最后优化:
[ \max_{\theta} \mathbb{E}{x, y \sim \pi{\theta_{\mathrm{old}}}} \left[ r_{\mathrm{SA}}(x, y)
\beta \log \frac{\pi_\theta(y|x)}{\pi_{\theta_{\mathrm{old}}}(y|x)} \right] ]
其中 (r_{\mathrm{SA}}(x, y)) 是自洽奖励,(\pi_{\theta_{\mathrm{old}}}) 是采样时冻结的旧策略,KL 项防止学生一次更新就把分布推到某个单一答案上。写成损失函数就是:
[ \mathcal{L}(\theta)
\frac{1}{K} \sum_{k=1}^{K} \hat{A}k \log \pi\theta(y_k|x) + \beta \cdot \mathrm{KL}\left(\pi_\theta(\cdot|x) | \pi_{\theta_{\mathrm{old}}}(\cdot|x)\right) ]
这里的 (\hat{A}_k) 是每个 prompt 内部归一化后的优势值:
[ \hat{A}k = c_k - \frac{1}{K} \sum{j=1}^K c_j ]
之所以要在每个 prompt 内部减均值,是因为不同 prompt 的自洽分数尺度不同。有的简单 prompt 采样十条全对,自洽分数都是 1;有的复杂 prompt 答案分散,最高权重才 0.4。如果直接用原始 reward,简单 prompt 的优势会盖过复杂 prompt,训练会被容易样本主导。
3.4 OPSA 和“蒸馏”的关系
从机制上说,OPSA 不是传统意义的知识蒸馏,因为没有一个外部教师分布需要被学生模仿。它更像“模型自己产生多条轨迹,再用自己的共识筛选轨迹,然后做策略更新”的自我蒸馏。如果把“蒸馏”宽泛地理解为“从一个分布提炼信号来训练另一个分布”,那么 OPSA 的教师可以看作“当前模型的多次采样投票分布”。
这一点恰好呼应了论文标题的提问:既然 on-policy 蒸馏本质上更像奖励优化,不如直接把目标改成明确的 on-policy 自对齐奖励,去掉“教师蒸馏”这层不准确的外壳。OPSA 的优势不是压缩模型,而是让模型在已有能力边界内,通过一致性信号提高稳定性和泛化性。
4. 最小实现:PyTorch 风格的 OPSA 训练循环
4.1 实现前需要确认的接口
OPSA 实现依赖三个核心能力:
- 模型能够批量生成多条候选序列;
- 能够重新计算每条序列在模型下的对数概率,用于策略优化;
- 需要有一个答案抽取和相等判断函数,用于计算自洽分数。
第一步先在数据集层面设计好extract_answer。如果是数学题,可以抽取最终数值;如果是选择题,抽取选项字母;如果是开放问答,需要先用文本向量或规则判断语义等价。这部分不严谨,后面的自洽训练会自动学到错误的奖励信号,所以宁可先把规则写慢,也不要跳过。
4.2 主要代码结构
下面的代码是 OPSA 训练循环的一种风格化实现,用于说明思路。实际项目需要根据模型库、数据格式和 tokenizer 细节调整。
import torch import torch.nn.functional as F from collections import Counter def batch_generate(model, tokenizer, prompts, k=8, max_new_tokens=512, temperature=0.8): outputs = [] for p in prompts: messages = tokenizer(p, return_tensors="pt") samples = [] for _ in range(k): ids = model.generate( **messages, max_new_tokens=max_new_tokens, do_sample=True, temperature=temperature, pad_token_id=tokenizer.eos_token_id, ) samples.append(ids[0][messages["input_ids"].shape[1]:]) outputs.append(samples) return outputs def answer_equality(a, b): # 不同任务替换成严格规则或语义近似 return a.strip() == b.strip() def compute_self_consistency(answers): counter = Counter(answers) scores = [counter[a] / len(answers) for a in answers] advantages = [s - sum(scores) / len(scores) for s in scores] return advantages def compute_seq_log_probs(model, input_ids, output_ids): # input_ids: (B, prompt_len) # output_ids: (K*B, seq_len) logits = model(input_ids=output_ids).logits shift_logits = logits[:, :-1, :].contiguous() shift_labels = output_ids[:, 1:].contiguous() log_probs = -F.cross_ent