跳转至

rob_ray_trainer.py — RobRayPPOTrainer 是机器人 VLA 模型的 PPO 训练器

文件路径: verl/experimental/vla/rob_ray_trainer.py 模块路径: verl.experimental.vla.rob_ray_trainer

文件概述

RobRayPPOTrainer 是机器人 VLA 模型的 PPO 训练器。它继承了通用的 PPO 训练框架,并为机器人场景做了大量定制:环境交互管理、轨迹收集与展平、异步环境管理等。这是 VLA PPO 训练的核心调度器。

核心类:RobRayPPOTrainer

初始化

class RobRayPPOTrainer:
    def __init__(self, config):
        self.config = config
        # 创建 Worker Groups
        self.actor_wg = ...     # Actor 模型 worker group
        self.critic_wg = ...    # Critic 模型 worker group
        self.env_wg = ...       # 环境 worker group
        self.rollout_wg = ...   # Rollout 推理 worker group

        # 创建 EnvLoop(环境交互循环)
        self.env_loop = EnvLoop(self.env_wg, self.rollout_wg, config)

训练循环 fit()

def fit(self):
    """主训练循环"""
    for step in range(self.total_steps):
        # 1. 环境重置
        self._reset_envs(step)

        # 2. 收集轨迹(通过 EnvLoop)
        trajectories = self.env_loop.rollout()

        # 3. 展平轨迹为训练数据
        train_data = self.flatten_trajectories(trajectories)

        # 4. 计算优势函数
        advantages = self._compute_advantages(train_data)

        # 5. 更新 Actor(PPO 策略梯度)
        self.actor_wg.update_policy(train_data)

        # 6. 更新 Critic(价值函数)
        self.critic_wg.update_value(train_data)

        # 7. 日志记录
        self._log_metrics(step)

环境重置策略

def _reset_envs(self, step):
    """重置环境到指定的初始状态

    支持从预定义的 state_ids 列表中选择初始状态,
    确保训练覆盖不同的任务和初始条件。
    """
    # 从数据集采样 state_ids 和 task_ids
    state_ids = self.dataset.sample_state_ids(batch_size)
    task_ids = self.dataset.sample_task_ids(batch_size)

    # 重置所有环境
    obs = self.env_wg.reset_envs_to_state_ids(state_ids, task_ids)

轨迹展平

def flatten_trajectories(self, trajectories):
    """将多步轨迹展平为单步转移

    机器人轨迹格式:
        [(obs_0, act_0, rew_0), (obs_1, act_1, rew_1), ...]

    展平后:
        独立的 (obs, act, rew, next_obs) 转移对
    """
    ...

Response Mask 计算

def compute_response_mask(self, data):
    """计算哪些 token 是模型的"响应"(动作预测)

    在 VLA 中,输入是图像+文本提示,响应是动作 token。
    Mask 标识出动作 token 的位置,用于 PPO 损失计算。
    """
    ...

核心类/函数列表

名称 类型 说明
RobRayPPOTrainer 类 机器人 PPO 训练器
fit 方法 主训练循环
_reset_envs 方法 环境重置(支持指定初始状态)
flatten_trajectories 方法 轨迹展平
compute_response_mask 方法 计算动作 token mask
_compute_advantages 方法 GAE 优势函数计算

与其他模块的关系

  • 使用 EnvLoop(env_loop.py)管理环境交互
  • 使用 RobActorRolloutRefWorker(fsdp_workers.py)管理 Actor 模型
  • 使用 EnvWorker(workers/env/env_worker.py)管理环境
  • 被 main_ppo.py 创建和调用
  • 数据来自 prepare_libero_dataset.py 生成的数据集

小结

RobRayPPOTrainer 是 VLA PPO 训练的"总指挥",协调模型推理、环境交互、奖励计算、策略更新等所有环节。它的特殊之处在于需要处理机器人环境的复杂性:多步轨迹、动态环境重置、异步环境管理等。