跳转至

utils.py — 该文件定义了 PPO 训练中的基础类型和辅助判断函数

文件概述

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

该文件定义了 PPO 训练中的基础类型和辅助判断函数。最重要的是 Role 枚举类,它定义了分布式训练中所有参与者的角色类型。

在训练流程中的位置

这是一个被广泛引用的基础模块,几乎所有训练相关的文件都依赖它来识别不同的 Worker 角色。

关键代码讲解

1. Role 枚举类

class Role(Enum):
    """角色枚举,定义分布式训练中的所有参与者角色"""
    Actor = 0           # 策略模型(仅作为 Actor)
    Rollout = 1         # 推理引擎(仅做生成)
    ActorRollout = 2    # Actor + Rollout 混合(最常用)
    Critic = 3          # 价值网络
    RefPolicy = 4       # 参考策略(用于 KL 惩罚)
    RewardModel = 5     # 奖励模型
    ActorRolloutRef = 6 # Actor + Rollout + Ref 三合一(LoRA 模式)
    Env = 7             # 环境(Agent 场景)

在实际使用中: - ActorRollout 是最常用的角色,因为 Actor 和 Rollout 通常共享同一个模型(混合引擎) - ActorRolloutRef 在使用 LoRA 时用到,因为 Ref Policy 就是去掉 LoRA adapter 的 Actor - RefPolicy 在全量微调(非 LoRA)时使用,需要独立的模型副本

2. Role 字符串转换

    def __str__(self):
        return self._get_role_string()

    def _get_role_string(self):
        role_mapping = {
            Role.Actor: "actor",
            Role.Rollout: "rollout",
            Role.ActorRollout: "actor_rollout",
            Role.Critic: "critic",
            Role.RefPolicy: "ref",
            Role.RewardModel: "rm",
            Role.ActorRolloutRef: "actor_rollout_ref",
        }
        return role_mapping.get(self, self.name.lower())

角色的字符串表示用于 checkpoint 路径命名(如 global_step_100/actor、global_step_100/critic)以及资源池映射。

3. 辅助判断函数

def need_reference_policy(config) -> bool:
    """判断是否需要参考策略"""
    return config.algorithm.use_kl_in_reward or config.actor_rollout_ref.actor.use_kl_loss

只有在启用 KL 奖励惩罚或 KL 损失时才需要参考策略。

def need_reward_model(config) -> bool:
    """判断是否需要奖励模型"""
    return config.reward.reward_model.enable
def need_critic(config) -> bool:
    """判断是否需要 Critic"""
    if config.critic.enable is not None:
        return bool(config.critic.enable)
    elif config.algorithm.adv_estimator == AdvantageEstimator.GAE:
        return True  # GAE 需要 Critic 提供 value 估计
    else:
        return False  # GRPO/REINFORCE++ 等不需要 Critic

这个函数的逻辑反映了不同算法对 Critic 的需求: - GAE(广义优势估计):需要 Critic 来估计状态价值 - GRPO/REINFORCE++:不需要 Critic,直接用奖励分数计算优势

核心类/函数列表

名称 类型 作用
Role Enum 定义训练中所有角色类型
WorkerType type alias type[Worker] 的别名
need_reference_policy(config) function 判断是否需要参考策略
need_reward_model(config) function 判断是否需要奖励模型
need_critic(config) function 判断是否需要 Critic 网络

数据流和调用关系

main_ppo.py
  +-- need_critic(config)           --> 决定是否创建 Critic Worker
  +-- need_reference_policy(config) --> 决定是否创建 Ref Policy Worker

ray_trainer.py
  +-- Role.ActorRollout   --> Worker Group 创建和管理
  +-- Role.Critic         --> Worker Group 创建和管理
  +-- Role.RefPolicy      --> Worker Group 创建和管理
  +-- Role.RewardModel    --> 资源池分配

小结

utils.py 是 PPO 训练的基石模块,提供了: 1. Role 枚举:统一定义所有角色类型,避免使用字符串带来的拼写错误 2. 判断函数:根据配置决定哪些组件需要被创建,实现了"按需创建"的灵活架构