on-policy distillation
定义
On-Policy Distillation是一种融合强化学习(On-Policy RL) 与知识蒸馏(Knowledge Distillation) 的模型训练范式,核心是让学生模型在自己生成的轨迹中学习,并由教师模型提供逐 token 密集监督,解决传统离线蒸馏的分布不匹配与RL反馈稀疏问题。
On-Policy(同策略):训练数据完全来自学生模型当前策略生成的轨迹(rollout),即 “学生自己做、自己学”。
Distillation(蒸馏):用一个更强的教师模型(Teacher) 指导更小 / 更弱的学生模型(Student),让学生对齐教师的输出分布。
On-Policy Distillation:学生先生成自己的输出序列,教师对这些序列逐 token 打分 / 给分布,学生再反向更新以缩小与教师的差异。
训练流程(三步闭环)
- 采样(Rollout):学生模型基于当前策略,对一批输入生成完整输出序列(轨迹)。
- 教师监督(Teacher Supervision):教师模型对学生生成的每个 token,输出其条件概率分布(逐 token 指导)。
- 更新(Update):学生最小化与教师在学生轨迹分布上的差异(常用反向 KL 散度),完成参数更新。
伪代码
def on_policy_distill_step(student, teacher, tokenizer):
# 1. 学生自己生成轨迹(on-policy data)
traj = generate_on_policy_batch(student, tokenizer, batch_size=rollout_batch_size)
mask = (traj != tokenizer.pad_token_id).long()
# 2. 把traj喂给学生 & 教师,得到traj上的logits
student_out = student(input_ids=traj, attention_mask=mask)
student_logits = student_out.logits[:, :-1, :] # 对齐 shift
with torch.no_grad():
teacher_out = teacher(input_ids=traj, attention_mask=mask)
teacher_logits = teacher_out.logits[:, :-1, :]
# 3. 标签与有效位置
labels = traj[:, 1:]
valid_mask = mask[:, 1:].reshape(-1)
# 4. On-policy 蒸馏损失:Reverse KL 散度
# Loss = E_student [ log π_student - log π_teacher ]
student_logp = F.log_softmax(student_logits, dim=-1)
teacher_logp = F.log_softmax(teacher_logits, dim=-1)
# 只在有效 token 上计算
loss = (
(student_logp - teacher_logp)
.gather(-1, labels.unsqueeze(-1))
.squeeze(-1)
* valid_mask
).mean()
return loss

浙公网安备 33010602011771号