0%

尝试了解 AgentGym-RL 并加入 OPD

fork 了 AgentGym-RL,并尝试加入 OPD 的 RFC。

AgentGym-RL OPD 适配

AgentGym-RL

AgentGym-RL 是基于 verl 的 trainer/worker/dataflow 框架,已有训练链路,在 verl 上增加了 multi-turn agent RL 等需要的代码。

训练的入口是 verl/agent_trainer/main_ppo.py。默认配置是用 Hydra 写的,在 verl/agent_trainer/config/ppo_trainer.yaml。用 Ray 分布式计算,主循环在 ray_trainer.pyfit() 中。

fit() 的伪代码:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
for epoch in total_epochs:
for batch in train_dataloader:
gen_batch = 从 batch 中取 prompt

gen_batch_output = actor_rollout_wg.generate_sequences(gen_batch) # 生成多轮轨迹

batch = batch.repeat(...) # 复制原始 batch 的元信息,让它和多条 rollout 对齐
batch = batch.union(gen_batch_output)

old_log_prob = actor_rollout_wg.compute_log_prob(batch) # 计算轨迹在模型下的 logprob

ref_log_prob = ref_policy_wg.compute_ref_log_prob(batch)

values = critic_wg.compute_values(batch) # 如果需要 critic

reward = batch["scores"]

batch = compute_advantage(batch) # 用 reward 算 advantage

critic_wg.update_critic(batch) # 如果需要 critic

actor_rollout_wg.update_actor(batch)

当调用 self.actor_rollout_wg.generate_sequences(gen_batch) 时,实际会进入 worker,在 verl/workers/agent_fsdp_workers.py,然后再进入 vLLM rollout,在 verl/workers/rollout/agent_vllm_rollout/vllm_rollout.py。得到的 output 是一整条 trajectory

worker 就是 Ray 的远程 worker。上述 RayPPOTrainer.fit() 是训练主控,它运行在 driver 进程。模型、GPU、torch 的 FSDP、vLLM engine 不在 driver 里,而在 ActorRolloutRefWorker 里。

verl/workers/rollout/schemas.py 里的 RolloutHandler 维护 loss_mask

verl/workers/agent_actor/dp_actor.py 里有 update_policy()

verl/agent_trainer/ppo/core_algos.py 里放 loss 函数。

OPD

对于 sampled-token on-policy distillation,学生模型生成动作 token,teacher 在同一批学生采样出来的 token 上给监督信号。

对一个生成出来的 token $a_t$,学生旧策略、学生当前策略、teacher 策略分别有:

一个轻量的 sampled-token OPD 信号可以写成:

直观理解,如果 teacher 比旧学生更喜欢这个 token,那么这个 token 的 teacher advantage 为正;如果 teacher 不喜欢,为负。

复用 PPO 风格的 clipped policy loss:

注意 mask,因为 AgentGym-RL 是多轮 agent 轨迹,有:

  • 初始 task prompt;
  • assistant action;
  • environment/user observation;
  • chat template token。

OPD 不能把 environment observation 当成学生 action 去蒸馏。本来可以考虑复用已有的 response_mask,它来自 rollout 阶段的 response_loss_mask,标记哪些 token 是训练目标。

但 AgentGym-RL 会在 rollout 后把 assistant 内容重新拼进 chat history,手工加入 assistant suffix,比如 Qwen 的 <|im_end|>。这个 suffix 不是学生模型真正 sampled 出来的 action token。对于 sampled-token OPD,更严格的做法是额外维护 action_mask,只蒸馏 student 实际生成的 assistant action/content token;如果某些 rollout 没有 action_mask,再 fallback 到 response_mask

AgentGym-RL OPD 适配

AgentGym-RL/verl/agent_trainer/ppo/core_algos.py 加入 OPD loss。

1
compute_sampled_token_distillation_loss(...)

默认 pg_reverse_kl:用 teacher logprob 和 old logprob 构造 teacher advantage,再复用 PPO clipped loss。

具体来说,对于输入张量中的三个核心:old_log_problog_probteacher_log_prob,先计算:

1
2
3
4
5
reverse_kl_est = verl_F.masked_mean(
old_log_prob - teacher_log_prob,
distillation_mask
)
teacher_advantages = teacher_log_prob - old_log_prob

然后复用 PPO clipped loss

1
2
3
4
5
6
7
distill_loss, distill_clipfrac, distill_kl = compute_policy_loss(
old_log_prob=old_log_prob,
log_prob=log_prob,
advantages=teacher_advantages,
eos_mask=distillation_mask,
cliprange=cliprange
)

AgentGym-RL/verl/agent_trainer/ppo/ray_trainer.py 加入 OPD 入口。trainer 是整条数据流的控制器。加了一个新角色 TeacherPolicy = 7。然后支持两种 teacher 来源:

  1. 如果配置了 algorithm.distillation.teacher_model.path,就启动一个单独的 teacher policy worker。
  2. 如果没有配置 teacher path,就复用已有的 ref_log_prob 作为 teacher_log_probs,用于便宜的 sanity check。

第二种用来验证数据流是否通了。

现在 batch 里多了一个字段,batch.batch["teacher_log_probs"]。数学上是 $\log \pi_T(a_t|s_t)$,teacher 对学生已经生成出来的 token 计算 logprob。

对于已经包含 teacher_log_probsbatch,主循环会执行

1
actor_output = self.actor_rollout_wg.update_actor(batch)

update_actor() 会进入 worker,然后进入 verl/workers/agent_actor/dp_actor.pyupdate_policy()

在这里,调用上述函数计算 OPD loss:

1
2
distill_loss, distill_metrics = core_algos.compute_sampled_token_distillation_loss(...)
policy_loss = policy_loss + distill_loss * loss_coef

就是

其中 loss_coef: 1.0 就是 $\lambda$。

这里默认使用 action_mask 作为 OPD 的 distillation_mask。所以它不会对 user/environment observation token 产生梯度,也不会蒸馏手工拼接的 assistant suffix token。如果 batch 里没有 action_mask,实现会退回到 response_mask

最后,AgentGym-RL/verl/agent_trainer/main_ppo.py 只在需要单独 teacher 时注册 Ray worker。这样默认情况下不会多起一个模型 worker,不会影响原来的训练。AgentGym-RL/verl/agent_trainer/config/ppo_trainer.yaml 默认关闭 OPD。加入示例脚本 examples/train/AgentGym-RL/searchqa_opd_train.sh,以及一个更小的 smoke 脚本 examples/train/AgentGym-RL/searchqa_opd_smoke.sh。加入测试 tests/test_opd_core_algos.py

测试

运行:

1
conda run -n agentgym-rl-test python -m pytest tests/test_opd_core_algos.py -q

通过。

语法检查:

1
2
3
4
5
6
conda run -n agentgym-rl-test python -m py_compile \
AgentGym-RL/verl/agent_trainer/ppo/core_algos.py \
AgentGym-RL/verl/workers/agent_actor/dp_actor.py \
AgentGym-RL/verl/agent_trainer/ppo/ray_trainer.py \
AgentGym-RL/verl/agent_trainer/main_ppo.py \
tests/test_opd_core_algos.py

通过。

启动 AgentGym environment server 真实试一下

待续,还没试。也许可以先用 examples/train/AgentGym-RL/searchqa_opd_smoke.sh 跑一个 0.5B/1 step 的数据流,验证 env server -> rollout -> old_log_probs -> ref-as-teacher log_probs -> OPD actor update

Reference

On-Policy Distillation - Thinking Machines Lab

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