思路:让 teacher 在 student 真实 rollout 的轨迹上去纠正分布.
普通蒸馏:
1for traj, y_teacher in offline_dataset: # x: (b, L)int. y_teacher (b, L)int.2 loss = KL(student(traj), teacher(traj)) # 一次性输出 len 个 dist. 需要一些 mask.OPD (on-policy distillation):
1for prompt in prompts: # x: (b, prompt_len)2 traj = student.generate(prompt) # (b, L)3 loss = KL(student(cat(prompt, traj)), teacher(cat(prompt, traj))) # 需要一些 mask.