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.py 的 fit() 中。
fit() 的伪代码:
1 | for epoch in total_epochs: |
当调用 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_prob、log_prob、teacher_log_prob,先计算:
1 | reverse_kl_est = verl_F.masked_mean( |
然后复用 PPO clipped loss:
1 | distill_loss, distill_clipfrac, distill_kl = compute_policy_loss( |
在 AgentGym-RL/verl/agent_trainer/ppo/ray_trainer.py 加入 OPD 入口。trainer 是整条数据流的控制器。加了一个新角色 TeacherPolicy = 7。然后支持两种 teacher 来源:
- 如果配置了
algorithm.distillation.teacher_model.path,就启动一个单独的 teacher policy worker。 - 如果没有配置 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_probs 的 batch,主循环会执行
1 | actor_output = self.actor_rollout_wg.update_actor(batch) |
update_actor() 会进入 worker,然后进入 verl/workers/agent_actor/dp_actor.py 的 update_policy()。
在这里,调用上述函数计算 OPD loss:
1 | distill_loss, distill_metrics = core_algos.compute_sampled_token_distillation_loss(...) |
就是
其中 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 | conda run -n agentgym-rl-test python -m py_compile \ |
通过。
启动 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