执行摘要
- 一句话:引入统一Trainer抽象,支持同步与异步PPO训练模式
- 推荐动作:### 建议
本 PR 是 PPO 训练器的重大架构重构,值得所有开发者和高级用户精读。重点关注抽象基类的设计模式、TransferQueue 集成方案以及异步模式的生命周期管理。但分离异步模式尚未完成,不建议生产环境使用。建议在后续 PR 中优先修复配置路径错误、完成 _compute_reward_colocate 实现并补充更多测试覆盖。
功能与动机
来自PR标题和body:'As title.',关联Issue为空。从代码变更可知,主要推动力是将原来独立的同步PPO训练器(main_ppo_sync.py)抽象化,使其能够统一支持同步、联合异步(co-located async)和分离异步(separate async)三种训练模式,并利用TransferQueue实现高效数据传输。此外,Issue评论指出联合异步模式需要vllm PR#44483的支持。
实现拆解
实现步骤
-
抽象基类PPOTrainer:在verl/trainer/ppo/v1/trainer_base.py中定义抽象基类,统一__init__设置必要组件(replay buffer、KL controller),声明init、_setup等抽象方法。原有main_ppo_sync.py的逻辑被抽取为基类实现,并引入register_trainer装饰器实现子类注册。
-
三种训练模式子类:
PPOTrainerSync:同步模式,在on_init_end中完成worker组和LLM server的初始化,训练循环内同步等待rollout完成。
PPOTrainerColocateAsync:联合异步模式,actor与rollout部署在同一GPU上,通过TransferQueue进行异步数据传输。
PPOTrainerSeparateAsync:分离异步模式,trainer与rollout分离部署,但当前__init__中抛出NotImplementedError,仅定义了骨架。此模式需要配置v1.separate_async.num_warmup_batches,但代码中错误地引用了colocate_async配置(review指出)。
-
ReplayBuffer与TransferQueue集成:新增verl/trainer/ppo/v1/replay_buffer.py,基于TransferQueue实现KV存储,按key格式{uid}_{session_id}_{index}管理轨迹,支持GRPO组采样控制(pending/running/finished/failure状态机)。
-
AgentLoopWorkerTQ适配器:新增verl/trainer/ppo/v1/agent_loop_tq.py,包装AgentLoopWorker以支持TransferQueue,实现generate_sequences中为每个样本创建后台异步任务(fire-and-forget),并通过_run_prompt控制采样参数。
-
入口重构:verl/trainer/main_ppo.py从直接定义TaskRunner改为加载TaskRunnerV1(通过load_class_from_fqn),原TaskRunner迁移至verl/trainer/main_ppo_v0.py作为向后兼容的V0实现。配置项trainer.use_v1控制选择哪条路径。
-
测试配套:新增tests/trainer/ppo/v2/test_replay_buffer_on_cpu.py(399行),模拟RolloutProducer在CPU上向ReplayBuffer写入数据,验证采样逻辑。
关键文件:
verl/trainer/ppo/v1/trainer_base.py(模块 训练器核心;类别 source;类型 rename-or-move;符号 PPOTrainer, compute_advantage_for_multi_trajectories, init, init): 抽象基类 PPOTrainer 的定义,重构核心训练逻辑,引入抽象方法 init/_setup 等。
verl/trainer/ppo/v1/agent_loop_tq.py(模块 代理循环;类别 source;类型 dependency-wiring;符号 apply_greedy_sampling_params, AgentLoopWorkerTQ, AgentLoopManagerTQ, generate_sequences): 新增 TransferQueue 适配的 AgentLoopWorkerTQ,支持异步生成序列,是关键数据路径。
verl/trainer/main_ppo_v0.py(模块 入口路由;类别 source;类型 dependency-wiring;符号 TaskRunner, init, add_actor_rollout_worker, add_critic_worker): V0 训练器的 Ray remote TaskRunner,提供向后兼容。
verl/trainer/ppo/v1/replay_buffer.py(模块 缓存层;类别 source;类型 dependency-wiring;符号 ReplayBuffer, init, _sync_metadata_from_transfer_queue, sample): 基于 TransferQueue 实现的 ReplayBuffer,替换原有 padding 机制,是异步数据管道的核心。
verl/trainer/ppo/v1/trainer_separate_async.py(模块 训练器分离异步;类别 source;类型 dependency-wiring;符号 HybridEngineMode, PPOTrainerSeparateAsync, init, _setup): 分离异步模式子类的骨架,当前抛出 NotImplementedError,但有独立的配置路径错误需修复。
verl/trainer/main_ppo.py(模块 主入口;类别 source;类型 dependency-wiring;符号 main, run_ppo, TaskRunner, TaskRunnerV1): 主入口重构,引入 TaskRunnerV1 路由,原有 TaskRunner 移至 main_ppo_v0.py。
tests/trainer/ppo/v2/test_replay_buffer_on_cpu.py(模块 测试;类别 test;类型 test-coverage;符号 tq_init, partition_id, _uid, _trajectory_key): CPU 上的 ReplayBuffer 单元测试,验证采样和状态机逻辑,新增 399 行测试代码。
关键符号:PPOTrainer, compute_advantage_for_multi_trajectories, init, _setup, ReplayBuffer, _poll_from_transfer_queue, close, apply_greedy_sampling_params, AgentLoopWorkerTQ, AgentLoopManagerTQ, generate_sequences, _run_prompt, _agent_loop_postprocess, create, TaskRunner, add_actor_rollout_worker, add_critic_worker, init_resource_pool_mgr, add_reward_model_resource_pool, add_teacher_model_resource_pool, add_ref_policy_worker, _sync_metadata_from_transfer_queue, sample, HybridEngineMode, PPOTrainerSeparateAsync, get_llm_client, on_init_end, on_train_begin, on_validate_begin, main, run_ppo, TaskRunnerV1, init_agent_loop_manager, tq_init, partition_id, _uid, _trajectory_key, PromptSpec, RolloutProducer, run
关键源码片段
verl/trainer/ppo/v1/trainer_base.py
抽象基类 PPOTrainer 的定义,重构核心训练逻辑,引入抽象方法 init/_setup 等。
# verl/trainer/ppo/v1/trainer_base.py
# 抽象基类 PPOTrainer,统一所有训练模式的初始化与生命周期
class PPOTrainer(ABC):
"""Base class for PPO trainer.
Args:
config: DictConfig from yaml config file.
"""
def __init__(self, config: DictConfig):
self.config = config
self.use_critic = need_critic(self.config)
self.use_reference_policy = need_reference_policy(self.config)
self.use_teacher_policy = need_teacher_policy(self.config)
if self.config.algorithm.use_kl_in_reward:
self.kl_ctrl_in_reward = core_algos.get_kl_controller(self.config.algorithm.kl_ctrl)
self.replay_buffer = ReplayBuffer() # 使用基于 TransferQueue 的 ReplayBuffer
def init(self):
"""Initialize all components of the trainer.
包括 WorkerGroup、LLMServerManager、CheckpointEngineManager、RewardLoopManager 等。
"""
raise NotImplementedError
verl/trainer/ppo/v1/agent_loop_tq.py
新增 TransferQueue 适配的 AgentLoopWorkerTQ,支持异步生成序列,是关键数据路径。
# verl/trainer/ppo/v1/agent_loop_tq.py
# TransferQueue 适配的 AgentLoopWorker,通过背景任务异步生成序列
@ray.remote
class AgentLoopWorkerTQ(AgentLoopWorker):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
tq.init()
self.background_tasks = set()
async def generate_sequences(self, batch: TensorDict) -> None:
"""为 batch 中每个样本启动 agent loop 任务,不等待结果"""
validate = batch.get("validate", False)
batch.pop("validate", None)
config = self.config.actor_rollout_ref.rollout
sampling_params = {
"temperature": config.temperature,
"top_p": config.top_p,
"top_k": config.top_k,
"repetition_penalty": 1.0,
"logprobs": config.calculate_log_probs,
}
if validate:
# 验证阶段使用 val_kwargs 覆盖采样参数
sampling_params.update({
"top_p": config.val_kwargs.top_p,
"top_k": config.val_kwargs.top_k,
"temperature": config.val_kwargs.temperature,
})
# 默认 single-turn agent
if "agent_name" not in batch:
batch["agent_name"] = NonTensorData(config.agent.default_agent_loop)
trajectory_info = await get_trajectory_info(batch["global_steps"], batch["index"], validate)
# 为每个样本创建 fire-and-forget 任务
for i in range(len(batch)):
prompt = {}
for k, v in batch.items():
if isinstance(v, torch.Tensor):
prompt[k] = v[i]
elif isinstance(v, NonTensorStack):
prompt[k] = v[i].data
elif isinstance(v, NonTensorData):
prompt[k] = v.data
else:
logger.error(f"Unsupported type {type(v)} for key {k}")
task = asyncio.create_task(
self._run_prompt(prompt, sampling_params, trajectory=trajectory_info[i])
)
self.background_tasks.add(task)
task.add_done_callback(self.background_tasks.discard)
verl/trainer/ppo/v1/replay_buffer.py
基于 TransferQueue 实现的 ReplayBuffer,替换原有 padding 机制,是异步数据管道的核心。
# verl/trainer/ppo/v1/replay_buffer.py
# 基于 TransferQueue 的轨迹缓存,支持 GRPO 组采样状态机
class ReplayBuffer:
def __init__(self, poll_interval: float = 2.0):
self.poll_interval = poll_interval
self.partitions: dict[str, dict[str, dict]] = defaultdict(dict)
self.pending_keys: dict[str, set] = defaultdict(set)
self.running_keys: dict[str, set] = defaultdict(set)
self.finished_keys: dict[str, set] = defaultdict(set)
self.failure_keys: dict[str, set] = defaultdict(set)
def _sync_metadata_from_transfer_queue(self):
"""从TransferQueue拉取所有key及其状态,更新内部状态集合"""
self.partitions.clear()
self.pending_keys.clear()
self.running_keys.clear()
self.finished_keys.clear()
self.failure_keys.clear()
data = tq.kv_list()
if data is None:
return
for partition_id, items in data.items():
partition = self.partitions[partition_id]
for key, tag in items.items():
if tag.get("is_prompt", False):
# GRPO group sampling 状态机:pending → running → finished/failure
match tag["status"]:
case "pending":
self.pending_keys[partition_id].add(key)
case "running":
self.running_keys[partition_id].add(key)
case "finished":
self.finished_keys[partition_id].add(key)
case "failure":
self.failure_keys[partition_id].add(key)
case _:
raise ValueError(f"Unknown status: {tag['status']}")
else:
# 普通轨迹数据
if key not in partition:
partition[key] = {}
partition[key].update(tag)
评论区精华
Review 讨论亮点
- 配置路径错误(高优先级):gemini-code-assist 指出
PPOTrainerSeparateAsync 中 num_warmup_batches 使用了错误的配置路径 colocate_async,应改为 separate_async。
- 方法签名不匹配(高优先级):gemini-code-assist 指出
_compute_reward_colocate 方法签名与实际调用不匹配,会引发 TypeError。作者回应该方法仍是 TODO 尚未实现。
-
logger.exception 误用(高优先级):gemini-code-assist 指出 agent_loop_tq.py 中在 except 块外使用 logger.exception 不当,应改用 logger.error。
-
配置路径错误:num_warmup_batches 使用了 colocate_async 而非 separate_async (correctness): 尚未修复,需在后续 PR 中更正。
- _compute_reward_colocate 签名与实际调用不匹配 (correctness): 作者回复该方法为 TODO 未实现,当前无影响,但未来需修复。
- agent_loop_tq.py 中 logger.exception 用在 except 块外 (style): 尚未修改。
风险与影响
关联脉络
- PR #6823 [BREAKING][trainer, cfg] chore: enable V1 trainer by default: 本 PR 是 V1 trainer 的实现基础,6823 在此基础上启用 V1 为默认模式。
- PR #6867 [fully_async, doc] fix: ignore temperature config for teacher prompt_logprobs and warn when non-default value is set: 与 fully_async 训练相关,可能与本 PR 的异步模式联动。
- PR #6882 [fully_async] fix: correct the use of partial_rollout: 涉及异步 rollout 修复,与本 PR 的异步模式有重叠。
参与讨论