数据截至 (上游 commit 352f1bd7c1a0)
04 · Trainer 与 verl:学习侧总装
本章讲什么: 前三章搭好了执行侧,这一章是学习侧怎么把它们串成训练循环:入口怎么起、rollout 怎么排怎么等、事件怎么变成训练样本、以及两个让「多轮 agent 轨迹」能被 GRPO 正确消化的关键定制。
1. 入口:站在 verl 的肩膀上
入口函数 run_ppo(agentlightning/verl/entrypoint.py:31-70)做三件事:初始化 Ray(把 agentlightning.verl.per_rollout_loss.register_in_worker 注册进每个 Ray actor 的启动钩子,:54-58)、把内存数据集包成 LoadedDataset(绕过 verl 的文件式数据集初始化,agentlightning/verl/dataset.py:17-31)、然后在 _AglTaskRunner 里组装 verl 的全套 worker(actor/critic/ref)——但训练器换成 AgentLightningRayPPOTrainer(entrypoint.py:130-144)。
AgentLightningRayPPOTrainer(agentlightning/verl/trainer.py:96)继承 verl 的 RayPPOTrainer:优化器、FSDP、KL、critic 这些重活全用 verl 原生的;被替换的是「rollout 怎么来」——原生 verl 在进程内让 vLLM 直接生成,这里改成「排队给真实 agent 跑」。每个训练步的主流程在 _train_step(trainer.py:444-648):
取训练批 ─▶ _rollout(走 Agent Lightning 采集)─▶ 转成 DataProto
─▶ 丢零优势样本 ─▶ 算优势 ─▶ 更新 critic/actor ─▶ 推权重到 vLLM
fit(trainer.py:650-736)就是标准 verl 训练循环:可选先验证、按 total_training_steps 迭代、周期性 _validate(走同一套 rollout 管线但不产训练批,trainer.py:738-771)和存 checkpoint。
2. Rollout 管理器:排队与等待
_rollout(trainer.py:363-442)每步先做两件准备:恢复 vLLM 副本的生成(_resume_all_rollout_generation,:159-161)、把 vLLM 地址重新注册进网关(register_model/delete_model,:378-397——先删后注,保证账本里没有上一轮的陈旧端点)。然后交给两种 rollout 管理器之一。
2.1 排队:GRPO 分组
_create_rollouts(agentlightning/verl/agl_rollout_manager.py:244-293)把训练批变成 rollout 队列:
- 每个样本生成一个
data_id(uuid)作为组键。 - 训练时每样本排
rollout.n个 rollout(GRPO 组大小,:250-257);验证时只排 1 个。 - 每个 rollout 预生成
rollout_id(:266-267)再配合post_with_retry——和第 01 章的幂等创建咬合,重试不产生重复。 - 排队前过一道
hooks.on_enqueue钩子(:263-264),用户可在此改写请求(如注入自定义 metadata)。
2.2 同步等待:完成即拉即删
AglRolloutManager.enqueue_and_wait_until_completed(agl_rollout_manager.py:394-444):每秒轮询每个未完成 rollout 的状态,一到终态就跑成功/失败钩子(:422-428)、拉取事件组装 CompletedRollout,并且立刻从账本删除(:409-412)——这就是第 01 章说「账本体积有界」的原因。挂起的 rollout 里谁先完成谁先处理,全批完成才返回。
2.3 异步等待:carry-over 组
AglAsyncRolloutManager.enqueue_and_wait_until_group_completed(agl_rollout_manager.py:450-533)是异步训练的引擎,逻辑变成「按组等」:GRPO 的优势是组内比较,必须同一 data_id 的 n 个 rollout 全部终态才算一组完成(:503-511)。每步目标凑够 train_batch_size 组;没凑完时,未完成的组整组结转(carry-over)到下一步(:527-533),下一步只补发差额(_next_train_batch_dict_for_rollout,trainer.py:267-275)。这让慢 agent 不阻塞训练节奏——最新的权重先训起来,慢样本晚几步照常进场。
启动异步模式有个硬约束:async_train_batch_size 必须大于 train_batch_size(trainer.py:103-109,否则没有结转余量, 直接 ValueError)。
2.4 事件 → Triplet
_build_completed_rollout(agl_rollout_manager.py:334-390)把一个完成 rollout 的事件流变成训练原料:
- 拉
format=triplet事件(第 01 章的精简通道)。 - 每个合法
model_request事件(无错误、有 response_token_ids)变成一个Triplet:prompt 侧token_ids、response 侧token_ids+log_probs(:338-359)。 - 最后一条
reward事件的标量作为final_reward,盖到最后一个 triplet 上(:361-368)——注意:v1.0 的信用分配就这么朴素,整趟一个分,配给最后一轮(组合/广播层面的细化在第 2 小节的 rollout 级优势里)。
生命周期时间戳(排队→运行→完成)也在这时从服务端状态里抠出来记好(_record_lifecycle_timestamps,:206-226,注释解释了动机:pod 分批启动导致 QUEUING 可以远晚于提交,取 updated_at 才是真实启动时刻),供训练侧输出排队/执行耗时的分位数指标(_rollout_lifecycle_metrics,trainer.py:209-265)。
3. RolloutAdapter:Triplet → verl 训练批
RolloutAdapter.get_train_data_batch(agentlightning/verl/rollout_adapter.py:211-502)把 Triplet 列表组装成 verl 的 DataProto。核心分叉是聚合粒度(trace_aggregator.level,默认 trajectory,agentlightning/verl/config.yaml:24-27):
3.1 transition 级:一轮模型调用 = 一行
每个 LLM 调用独立成一行训练样本,同一个 rollout 的每一轮都配上同一个 final_reward(rollout_adapter.py:316-330)。简单直接,但一个 10 轮 agent 会膨胀成 10 行,每行都要单独过一遍前向。
3.2 trajectory 级:整条轨迹拼成一行(默认)
把多轮调用按 token 前缀连续性拼成尽量少的行(rollout_adapter.py:331-418)。拼接判据是 ids_startswith(:45):下一轮 prompt 的 token 序列必须精确地以上一轮累积上下文为前缀(多轮对话模板保证这一点)。拼的时候有个精妙处理——工具观察段(observation):
# 示意,对应 rollout_adapter.py:349-355 的合并逻辑
if ids_startswith(prompt_ids, current_context):
observation = prompt_ids[len(current_context):] # 前缀之外的新增部分 = 工具结果
current_response_ids += observation # 拼进 response(上下文连续)
current_response_mask += [0] * len(observation) # 但 mask=0:不参与训练
即:agent 轨迹「模型说 → 工具回 → 模型说 → …」被拼成一条连续序列,模型生成的段 mask=1 训练,工具返回的段 mask=0 只当上下文。这样一条轨迹一次前向就完成,多轮 agent 的训练成本大幅下降。前缀对不上的轮次则断开另起一行,并把两个版本的解码文本记进 wandb 表供人工排查(merge_mismatch_rows,:366-384)。响应侧 logprobs 同步拼接;任何一轮缺失或长度不齐,整行的 logprobs 弃用(:354-362, 272-273)。
3.3 组装成张量
行级细节(append_training_row,:246-303):prompt 超 max_prompt_length 截断并打 is_drop 标记;response 超长截断;prompt 左 pad、response 右 pad(get_left_padded_ids_and_attention_mask/get_right_padded_ids_and_attention_mask,:167-186——verl 的训练约定)。奖励放在每行最后一个有效 token 位置上(token_level_scores,:451-455)——verl 的 outcome reward 惯例。若所有行都带采样 logprobs,则整批附带 rollout_log_probs(:443-475),供 bypass mode/rollout correction 使用(§4.2)。
4. 两个默认开启的定制
4.1 rollout 级优势(rollout-level advantage)
GRPO 原生按「行」算组内优势;但 trajectory 聚合后一个 rollout 可能拆成多行,组结构被破坏。compute_rollout_level_advantage(agentlightning/verl/rollout_level_advantage.py:15-99,默认开启:config.yaml:10 enable_rollout_level_advantage: true)的做法:每个 rollout 取一行代表(校验同 rollout 各行 uid 与奖励和一致,:116-128),用 verl 原生的 compute_advantage 在代表行上算 GRPO 组内优势,再把标量优势广播回该 rollout 的所有行所有训练 token(_broadcast_rollout_scalars,:154-162)——同一次执行的每一段共享同一个优势值,组结构以 rollout 为单位重建。
4.2 逐 rollout 均值损失(per_rollout_mean)
verl 默认的 token-mean 损失让 token 多的 rollout 天然贡献更大梯度。v1.0 默认换成 per_rollout_mean(config.yaml:38-39 loss_mode: per_rollout_mean):先把每行优势除以「该 rollout 总 token 数 × 批内行数」(normalize_advantages_by_rollout,agentlightning/verl/per_rollout_loss.py:14-39),再用注册进 verl 的自定义 policy loss(compute_policy_loss_per_rollout_mean,:42-87,带双向 clip_ratio 和 clip_ratio_c 的完整 PPO clip 实现)聚合——每个 rollout 对梯度的贡献均等,长轨迹不会主导训练。注册机制用的正是入口处挂的 Ray worker hook(第 1 节)。
5. 采集与训练的切换:暂停排空与零优势丢弃
暂停排空:异步模式下 rollout 结束、训练开始前,_pause_and_drain_gateway(trainer.py:167-189)暂停网关(pause 请求本身给 300 秒超时,防在途请求打满服务端)、然后每 0.25 秒轮询 inflight 直到归零,接着才让 vLLM 副本睡眠(checkpoint_manager.sleep_replicas(),trainer.py:539)、算优势、更新权重,下一步开始再恢复网关与生成(trainer.py:371-373)。同步模式则简单粗暴:直接 abort vLLM 上所有残余请求(_abort_all_rollout_requests,:155-157)。两套协议共同保证:权重切换瞬间没有跨版本的在途生成。
零优势丢弃:mini-batch 对齐需要丢掉多余样本时,优先丢「同组奖励全相同」的样本(_same_reward_uid_indices + 丢弃逻辑,:78-93, 511-530)——GRPO 组内无奖励方差则优势恒为 0,训了也没梯度,丢它们最不心疼;指标 training/n_zero_adv_groups 同步输出(_grpo_group_metrics,:61-75)。
REMAX 不支持:显式 NotImplementedError(:479-480)。
6. 生命周期钩子(hooks)
RolloutHooks(agentlightning/hooks.py:19-33)四个可覆盖点:on_startup、on_enqueue(排队前改写请求)、on_succeeded/on_failed(终态后回看全部事件并可通过 TraceWriter 回写自定义事件)。从单文件加载、要求恰好一个 RolloutHooks 子类(load_hooks,:36-63)。钩子事件先进 _TraceEventHelper 缓冲、钩子跑完统一 flush(agl_rollout_manager.py:91-106)——钩子里抛异常只打印不中断训练(:317-321)。
小结: 学习侧的全部巧思在于「让 verl 吃下多轮 agent 轨迹」:GRPO 以 data_id 分组、以组为完成单位(异步还能结转);trajectory 聚合把多轮拼成一行并用 mask 排除工具观察;rollout 级优势和逐 rollout 损失修正了轨迹长短带来的梯度偏差;暂停排空守住权重切换的一致性。最后一章收口全书。