0%

OPD 极简笔记

OPD 极简笔记,方便回顾。

回顾

记号

SFT

Hard-label SFT:

Soft-label / distillation SFT:

RL

目标:

Policy gradient:

GRPO

同一 prompt 采样多个 answers:

reward:

Group-relative advantage:

Ratio:

PPO-style clipped loss:

OPD

目标

OPD 目标:student-prefix 上的 reverse KL。

固定一个 prefix:

Full-vocabulary OPD

Full-vocabulary loss:

Gradient:

Sampled-token OPD

Student rollout 采一个 token:

Sampled-token reverse KL:

OPD advantage:

Surrogate loss

真正目标:

RL 实际训练用 surrogate。

最简单 PG surrogate:

Importance-sampling ratio:

IS surrogate loss:

把 OPD advantage 塞进去:

得到最终 OPD-IS loss:

基于 Tinker 的代码实现

通用 Tinker importance-sampling loss 需要:

1
2
3
4
new_logprobs      # log π_θ(a_t | s_t)
old_logprobs # log π_old(a_t | s_t)
advantages # A_t
mask # valid token mask

importance_sampling_loss:

1
2
3
4
def importance_sampling_loss(new_logprobs, old_logprobs, advantages, mask):
ratio = torch.exp(new_logprobs - old_logprobs)
loss = -(ratio * advantages * mask).sum() / mask.sum()
return loss

OPD 只负责制造 advantages:

1
2
reverse_kl = old_logprobs - teacher_logprobs
advantages = -reverse_kl

等价于:

1
advantages = teacher_logprobs - old_logprobs

然后传给已有 RL loss:

1
2
3
4
5
6
loss = importance_sampling_loss(
new_logprobs=new_logprobs,
old_logprobs=old_logprobs,
advantages=advantages,
mask=mask,
)

sampled-token OPD 如下:

1
2
3
4
5
6
7
8
9
10
sampled_logprobs = trajectories.loss_fn_inputs["logprobs"]
teacher_logprobs = teacher_client.compute_logprobs(trajectories)

reverse_kl = sampled_logprobs - teacher_logprobs
advantages = -reverse_kl

training_client.forward_backward(
trajectories,
loss_fn="importance_sampling",
)

直观看一下梯度

令:

单 token loss:

求导:

梯度下降:

因此:

Reference

On-Policy Distillation - Thinking Machines Lab

thinking-machines-lab/tinker-cookbook 的 on_policy_distillation.py