跳转至

ray_trainer.py — 这是 整个 verl 框架最核心的文件

文件概述

模块路径: verl.trainer.ppo.ray_trainer

这是 整个 verl 框架最核心的文件,实现了基于 Ray 的分布式 PPO 训练器 RayPPOTrainer。它作为"单控制器"(single controller),在一个驱动进程上协调多个分布式 Worker Group,完成 PPO 训练的完整流程。

该文件约 1600 行代码,包含: - KL 惩罚计算 - 优势估计 - 完整的分布式训练循环(fit()) - Worker Group 初始化 - Checkpoint 保存/加载 - 验证和指标收集

在训练流程中的位置

这是训练流程的"心脏",由 main_ppo.py 的 TaskRunner.run() 创建并调用 fit() 启动训练。

关键代码讲解

1. 顶层辅助函数

KL 惩罚

def apply_kl_penalty(data: DataProto, kl_ctrl, kl_penalty="kl"):
    """在 token 级别的奖励上施加 KL 惩罚"""
    # 计算当前策略和参考策略之间的 KL 散度
    kld = core_algos.kl_penalty(
        data.batch["old_log_probs"],
        data.batch["ref_log_prob"],
        kl_penalty=kl_penalty
    )
    kld = kld * response_mask
    beta = kl_ctrl.value  # KL 系数

    # 原始奖励 - KL 惩罚
    token_level_rewards = token_level_scores - beta * kld

    # 更新 KL 控制器(自适应模式下会调整 beta)
    kl_ctrl.update(current_kl=current_kl, n_steps=batch_size)
    data.batch["token_level_rewards"] = token_level_rewards

KL 惩罚的作用是防止策略偏离参考策略太远,是 RLHF 中的关键技术。

优势估计入口

def compute_advantage(data, adv_estimator, gamma=1.0, lam=1.0, ...):
    """根据选择的优势估计器计算优势值"""
    if adv_estimator == AdvantageEstimator.GAE:
        advantages, returns = core_algos.compute_gae_advantage_return(
            token_level_rewards=data.batch["token_level_rewards"],
            values=data.batch["values"],        # 来自 Critic
            response_mask=data.batch["response_mask"],
            gamma=gamma, lam=lam,
        )
    elif adv_estimator == AdvantageEstimator.GRPO:
        advantages, returns = core_algos.compute_grpo_outcome_advantage(
            token_level_rewards=data.batch["token_level_rewards"],
            response_mask=grpo_calculation_mask,
            index=data.non_tensor_batch["uid"],  # 按 prompt 分组
            ...
        )
    else:
        # 通过注册表动态获取(REINFORCE++, RLOO, OPO 等)
        adv_estimator_fn = core_algos.get_adv_estimator_fn(adv_estimator)
        advantages, returns = adv_estimator_fn(**adv_kwargs)

    data.batch["advantages"] = advantages
    data.batch["returns"] = returns

这个函数根据配置选择不同的优势估计方法。优势估计是 RL 训练质量的关键。

2. RayPPOTrainer 类

初始化

class RayPPOTrainer:
    def __init__(self, config, tokenizer, role_worker_mapping, resource_pool_manager, ...):
        self.tokenizer = tokenizer
        self.config = config
        self.hybrid_engine = config.actor_rollout_ref.hybrid_engine
        assert self.hybrid_engine, "Currently, only support hybrid engine"

        # 判断需要哪些组件
        self.use_reference_policy = need_reference_policy(self.config)
        self.use_rm = need_reward_model(self.config)
        self.use_critic = need_critic(self.config)

        # LoRA 模式下,Ref Policy 就是没有 LoRA 的 Actor
        self.ref_in_actor = lora_rank > 0 or config.actor_rollout_ref.model.get("lora_adapter_path") is not None

        # KL 控制器
        if self.config.algorithm.use_kl_in_reward:
            self.kl_ctrl_in_reward = core_algos.get_kl_controller(self.config.algorithm.kl_ctrl)

        self._create_dataloader(train_dataset, val_dataset, collate_fn, train_sampler)

init_workers() - Worker Group 创建

    def init_workers(self):
        """创建分布式训练的 Worker Group"""
        # 1. 创建资源池
        self.resource_pool_manager.create_resource_pool()

        # 2. 为每个角色创建 RayClassWithInitArgs
        actor_rollout_cls = RayClassWithInitArgs(
            cls=self.role_worker_mapping[actor_role],
            config=self.config.actor_rollout_ref,
            role=str(actor_role),
        )

        # 3. 创建 Critic Worker(如果需要)
        if self.use_critic:
            critic_cls = RayClassWithInitArgs(cls=self.role_worker_mapping[Role.Critic], ...)

        # 4. 使用 create_colocated_worker_cls 创建共置 Worker
        #    (多个角色共享同一组 GPU)
        for resource_pool, class_dict in self.resource_pool_to_cls.items():
            worker_dict_cls = create_colocated_worker_cls(class_dict=class_dict)
            wg_dict = self.ray_worker_group_cls(
                resource_pool=resource_pool,
                ray_cls_with_init=worker_dict_cls,
            )
            spawn_wg = wg_dict.spawn(prefix_set=class_dict.keys())
            all_wg.update(spawn_wg)

        # 5. 提取各角色的 Worker Group
        self.actor_rollout_wg = all_wg[str(actor_role)]
        self.actor_rollout_wg.init_model()
        if self.use_critic:
            self.critic_wg = all_wg[str(Role.Critic)]
            self.critic_wg.init_model()

        # 6. 创建 Reward Loop Manager 和 Checkpoint Manager
        self.reward_loop_manager = RewardLoopManager(config=self.config, ...)
        self.async_rollout_manager = AgentLoopManager.create(...)
        self.checkpoint_manager = CheckpointEngineManager(...)

这里的 create_colocated_worker_cls 是关键设计:它允许多个角色(如 Actor 和 Critic)共享同一组 GPU,在需要时切换模型,以节省 GPU 内存。

3. fit() - 核心训练循环

    def fit(self):
        """PPO 训练主循环"""
        logger = Tracking(...)
        self.global_steps = 0

        # 加载 checkpoint 并同步权重
        self._load_checkpoint()
        self.checkpoint_manager.update_weights(self.global_steps)

        # 训练前验证
        if self.config.trainer.get("val_before_train", True):
            val_metrics = self._validate()
            logger.log(data=val_metrics, step=self.global_steps)

        for epoch in range(current_epoch, self.config.trainer.total_epochs):
            for batch_dict in self.train_dataloader:
                metrics = {}
                timing_raw = {}
                batch = DataProto.from_single_dict(batch_dict)

Step 1: 生成回复

                with marked_timer("gen", timing_raw):
                    gen_batch_output = self.async_rollout_manager.generate_sequences(gen_batch_output)
                    self.checkpoint_manager.sleep_replicas()

调用推理引擎生成回复。生成完成后,让推理副本休眠以释放 GPU 显存。

Step 2: 计算奖励

                with marked_timer("reward", timing_raw):
                    if self.use_rm and "rm_scores" not in batch.batch.keys():
                        batch_reward = self._compute_reward_colocate(batch)
                        batch = batch.union(batch_reward)
                    reward_tensor, reward_extra_infos_dict = extract_reward(batch)

Step 3: 计算旧策略概率

                with marked_timer("old_log_prob", timing_raw):
                    old_log_prob, old_log_prob_mfu = self._compute_old_log_prob(batch)
                    # 计算熵(用于正则化)
                    entropys = old_log_prob.batch["entropys"]
                    batch = batch.union(old_log_prob)

Step 4: 计算参考策略概率(如果需要 KL 惩罚)

                if self.use_reference_policy:
                    with marked_timer(str(Role.RefPolicy), timing_raw):
                        ref_log_prob = self._compute_ref_log_prob(batch)
                        batch = batch.union(ref_log_prob)

Step 5: 计算价值函数(如果使用 Critic)

                if self.use_critic:
                    with marked_timer("values", timing_raw):
                        values = self._compute_values(batch)
                        batch = batch.union(values)

Step 6: 计算优势和奖励

                with marked_timer("adv", timing_raw):
                    batch.batch["token_level_scores"] = reward_tensor
                    if self.config.algorithm.use_kl_in_reward:
                        batch, kl_metrics = apply_kl_penalty(batch, ...)
                    else:
                        batch.batch["token_level_rewards"] = batch.batch["token_level_scores"]

                    # Rollout 校正(处理 off-policy 问题)
                    if rollout_corr_config is not None and "rollout_log_probs" in batch.batch:
                        batch, is_metrics = compute_rollout_correction_and_add_to_batch(batch, ...)

                    # 计算优势
                    batch = compute_advantage(batch, adv_estimator=..., gamma=..., lam=...)

Step 7: 更新 Critic 和 Actor

                if self.use_critic:
                    with marked_timer("update_critic", timing_raw):
                        critic_output = self._update_critic(batch)

                if self.config.trainer.critic_warmup <= self.global_steps:
                    with marked_timer("update_actor", timing_raw):
                        actor_output = self._update_actor(batch)

critic_warmup 允许 Critic 先训练几步再开始更新 Actor,使得价值估计更稳定。

Step 8: 同步权重、保存和验证

                    # 保存 checkpoint
                    if self.config.trainer.save_freq > 0 and (is_last_step or ...):
                        self._save_checkpoint()

                    # 同步权重到推理引擎
                    with marked_timer("update_weights", timing_raw):
                        self.checkpoint_manager.update_weights(self.global_steps)

                # 验证
                if self.config.trainer.test_freq > 0 and (is_last_step or ...):
                    val_metrics = self._validate()
                    metrics.update(val_metrics)

                # 记录所有指标
                metrics.update(compute_data_metrics(batch=batch, use_critic=self.use_critic))
                metrics.update(compute_timing_metrics(batch=batch, timing_raw=timing_raw))
                metrics.update(compute_throughout_metrics(batch=batch, timing_raw=timing_raw, n_gpus=n_gpus))
                logger.log(data=metrics, step=self.global_steps)

4. 序列长度均衡

    def _balance_batch(self, batch, metrics, ...):
        """重新排列数据,使每个 DP rank 获得相似的总 token 数"""
        global_seqlen_lst = batch.batch["attention_mask"].view(batch_size, -1).sum(-1)
        dp_size = self._get_dp_size(self.actor_rollout_wg, "actor")
        global_partition_lst = get_seqlen_balanced_partitions(workload_lst, k_partitions=dp_size, ...)
        batch.reorder(global_idx)

这是性能优化的关键:如果不均衡,某些 rank 的序列特别长,其他 rank 就要等待,造成 GPU 空闲。

核心类/函数列表

名称 类型 作用
apply_kl_penalty() function 在奖励上施加 KL 惩罚
compute_response_mask() function 计算回复部分的 attention mask
compute_advantage() function 根据配置选择优势估计方法
RayPPOTrainer class 核心训练器类
RayPPOTrainer.__init__() method 初始化配置、数据加载器等
RayPPOTrainer.init_workers() method 创建分布式 Worker Group
RayPPOTrainer.fit() method 执行完整训练循环
RayPPOTrainer._compute_values() method 调用 Critic 计算价值函数
RayPPOTrainer._compute_ref_log_prob() method 调用 Ref Policy 计算参考概率
RayPPOTrainer._compute_old_log_prob() method 调用 Actor 计算旧策略概率
RayPPOTrainer._update_actor() method 更新 Actor 策略网络
RayPPOTrainer._update_critic() method 更新 Critic 价值网络
RayPPOTrainer._save_checkpoint() method 保存 checkpoint
RayPPOTrainer._load_checkpoint() method 加载 checkpoint
RayPPOTrainer._validate() method 执行验证
RayPPOTrainer._balance_batch() method 序列长度均衡

数据流和调用关系

fit()
  |
  +-- _load_checkpoint()
  +-- checkpoint_manager.update_weights()
  +-- _validate()  (训练前验证)
  |
  +-- for epoch:
        for batch in train_dataloader:
          |
          +-- generate_sequences()   --> gen_batch_output (含 responses, rollout_log_probs)
          |
          +-- _compute_reward()      --> reward_tensor (rm_scores)
          |
          +-- _compute_old_log_prob() --> old_log_probs, entropys
          |
          +-- _compute_ref_log_prob() --> ref_log_prob (如果需要 KL)
          |
          +-- _compute_values()       --> values (如果使用 Critic)
          |
          +-- apply_kl_penalty()      --> token_level_rewards
          |
          +-- compute_advantage()     --> advantages, returns
          |
          +-- _update_critic()        --> 更新 Critic
          |
          +-- _update_actor()         --> 更新 Actor
          |
          +-- checkpoint_manager.update_weights()  --> 同步到推理引擎
          |
          +-- _save_checkpoint()      (定期)
          +-- _validate()             (定期)
          +-- logger.log(metrics)

小结

ray_trainer.py 是 verl 框架的核心,它实现了:

  1. 完整的 PPO 训练循环:从生成回复到计算奖励、优势估计、策略更新
  2. 分布式协调:通过 Ray Worker Group 管理多个分布式组件
  3. 灵活的算法支持:支持 GAE、GRPO、REINFORCE++ 等多种优势估计
  4. 高效的资源管理:共置模型、序列长度均衡、休眠/唤醒机制
  5. 完善的工程特性:checkpoint、断点恢复、指标收集、性能分析

理解这个文件是理解整个 verl 训练系统的关键。