on-policy蒸馏是伪蒸馏?OPSA自对齐重塑大模型训练

📅 发布时间:2026/9/5 23:47:30
on-policy蒸馏是伪蒸馏?OPSA自对齐重塑大模型训练 知识蒸馏进入大模型时代后一个研究方法论层面的问题被摆到了台面上当教师模型不再作为离线标签生成器而是顺着学生模型自己采样的分布逐 token 给出概率目标这种 on-policy 蒸馏到底是在“蒸馏”还是在做另一种策略强化更新围绕这个提问相关论文会重新拆解 on-policy 蒸馏的优化目标并给出不依赖外部教师监督的 OPSAOn-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-PolicyP 是 PolicyS 是 SelfA 是 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_{k1}^{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{j1}^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, k8, max_new_tokens512, temperature0.8): outputs [] for p in prompts: messages tokenizer(p, return_tensorspt) samples [] for _ in range(k): ids model.generate( **messages, max_new_tokensmax_new_tokens, do_sampleTrue, temperaturetemperature, pad_token_idtokenizer.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_idsoutput_ids).logits shift_logits logits[:, :-1, :].contiguous() shift_labels output_ids[:, 1:].contiguous() log_probs -F.cross_ent