Prhub

#6710 [trainer] feat: add unify trainer abstraction for sync and async training

原始 PR 作者 wuxibin89 合并时间 2026-06-12 22:43 文件变更 27 提交数 9 评论 6 代码增减 +2037 / -1151

执行摘要

引入统一 Trainer 抽象,支持同步与异步 PPO 训练模式

来自PR标题和body:'As title.',关联Issue为空。从代码变更可知,主要推动力是将原来独立的同步PPO训练器(main_ppo_sync.py)抽象化,使其能够统一支持同步、联合异步(co-located async)和分离异步(separate async)三种训练模式,并利用TransferQueue实现高效数据传输。此外,Issue评论指出联合异步模式需要vllm PR#44483的支持。

建议

本 PR 是 PPO 训练器的重大架构重构,值得所有开发者和高级用户精读。重点关注抽象基类的设计模式、TransferQueue 集成方案以及异步模式的生命周期管理。但分离异步模式尚未完成,不建议生产环境使用。建议在后续 PR 中优先修复配置路径错误、完成 _compute_reward_colocate 实现并补充更多测试覆盖。

讨论亮点

Review 讨论亮点

  • 配置路径错误(高优先级):gemini-code-assist 指出 PPOTrainerSeparateAsyncnum_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

实现拆解

实现步骤

  1. 抽象基类PPOTrainer:在verl/trainer/ppo/v1/trainer_base.py中定义抽象基类,统一__init__设置必要组件(replay buffer、KL controller),声明init_setup等抽象方法。原有main_ppo_sync.py的逻辑被抽取为基类实现,并引入register_trainer装饰器实现子类注册。

  2. 三种训练模式子类

    • 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指出)。
  3. ReplayBuffer与TransferQueue集成:新增verl/trainer/ppo/v1/replay_buffer.py,基于TransferQueue实现KV存储,按key格式{uid}_{session_id}_{index}管理轨迹,支持GRPO组采样控制(pending/running/finished/failure状态机)。

  4. AgentLoopWorkerTQ适配器:新增verl/trainer/ppo/v1/agent_loop_tq.py,包装AgentLoopWorker以支持TransferQueue,实现generate_sequences中为每个样本创建后台异步任务(fire-and-forget),并通过_run_prompt控制采样参数。

  5. 入口重构verl/trainer/main_ppo.py从直接定义TaskRunner改为加载TaskRunnerV1(通过load_class_from_fqn),原TaskRunner迁移至verl/trainer/main_ppo_v0.py作为向后兼容的V0实现。配置项trainer.use_v1控制选择哪条路径。

  6. 测试配套:新增tests/trainer/ppo/v2/test_replay_buffer_on_cpu.py(399行),模拟RolloutProducer在CPU上向ReplayBuffer写入数据,验证采样逻辑。

文件 模块 状态 重要度
verl/trainer/ppo/v1/trainer_base.py 训练器核心 renamed 9.25
verl/trainer/ppo/v1/agent_loop_tq.py 代理循环 added 9.23
verl/trainer/main_ppo_v0.py 入口路由 added 9.22
verl/trainer/ppo/v1/replay_buffer.py 缓存层 added 9.02
verl/trainer/ppo/v1/trainer_separate_async.py 训练器分离异步 added 9.02
verl/trainer/main_ppo.py 主入口 modified 8.93
tests/trainer/ppo/v2/test_replay_buffer_on_cpu.py 测试 added 8.14

关键符号

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 rename-or-move

抽象基类 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 dependency-wiring

新增 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 dependency-wiring

基于 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)

评论区精华

配置路径错误:num_warmup_batches 使用了 colocate_async 而非 separate_async 正确性

gemini-code-assist 指出 PPOTrainerSeparateAsync 中 num_warmup_batches 应从 self.config.trainer.v1.separate_async.num_warmup_batches 获取,但代码误用了 colocate_async 路径。

结论:尚未修复,需在后续 PR 中更正。 · 待处理

_compute_reward_colocate 签名与实际调用不匹配 正确性

gemini-code-assist 指出 method 定义为 def _compute_reward_colocate(self, batch, metrics),但调用处只传 batch 一个参数,将引发 TypeError。

结论:作者回复该方法为 TODO 未实现,当前无影响,但未来需修复。 · 已标记 TODO

agent_loop_tq.py 中 logger.exception 用在 except 块外 style

gemini-code-assist 指出 logger.exception 应在 except 块内使用,否则会捕获异常栈不清,建议改用 logger.error。

结论:尚未修改。 · 待处理

风险与影响

风险分析

  • 分离异步模式未实现PPOTrainerSeparateAsync.__init__ 抛出 NotImplementedError,若用户尝试使用该模式将直接崩溃。
  • 配置路径错误trainer_separate_async.pynum_warmup_batches 错误引用 colocate_async 配置,可能导致使用错误值。
  • _compute_reward_colocate 签名问题:虽然当前为 TODO,但未来若调用会触发 TypeError
  • TransferQueue 硬依赖:去除了原来的 try-except 降级逻辑,环境中未安装 TransferQueue 将导致启动失败。
  • V0/V1 兼容性main_ppo.py 路径选择依赖新配置项,旧脚本若未包含 use_v1 可能意外路由到 V1。

影响分析

  • 用户影响:需配置 trainer.use_v1trainer.v1.trainer_mode 来选择训练模式。旧配置若未更新可能触发断言或使用默认值。
  • 系统影响:训练器架构统一,后续异步训练和实验性功能开发更便捷。但部分异步路径尚未完善。
  • 团队影响:核心训练逻辑从单体文件拆分为多文件,降低认知负荷,但新架构需要团队学习。
  • 兼容性main_ppo.py 入口保留,原有依赖 RayPPOTrainer 的代码可通过 V0 继续运行。
分离异步模式未实现 配置路径错误 logger.exception 误用 TransferQueue 硬依赖

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论