跳转至

env_worker.py — EnvWorker 是运行在 Ray 上的环境工作器

文件路径: verl/experimental/vla/workers/env/env_worker.py 模块路径: verl.experimental.vla.workers.env.env_worker

文件概述

EnvWorker 是运行在 Ray 上的环境工作器,管理多个 EnvManager 实例(每个对应一个流水线阶段)。它是 verl 单控制器架构的一部分,通过 @register 装饰器暴露方法给远程调用。

核心类:EnvWorker

初始化

class EnvWorker(Worker, DistProfilerExtension):
    def __init__(self, config: DictConfig):
        Worker.__init__(self)
        self.stage_num = self.cfg.rollout.pipeline_stage_num

        # 初始化分布式
        initialize_global_process_group_ray(timeout_second=None)
        env_device_mesh = init_device_mesh(device_name, mesh_shape=(self.world_size, 1))

        # 每个阶段一个模拟器管理器
        self.simulator_list = []

初始化环境

@register(dispatch_mode=Dispatch.ONE_TO_ALL)
def init_worker(self):
    """根据配置创建环境"""
    if self.cfg.train.simulator_type == "libero":
        from verl.experimental.vla.envs.libero_env.libero_env import LiberoEnv
        for _ in range(self.stage_num):
            self.simulator_list.append(
                EnvManager(self.cfg.train, rank=self._rank,
                          world_size=self._world_size, env_cls=LiberoEnv)
            )
    elif self.cfg.train.simulator_type == "isaac":
        from verl.experimental.vla.envs.isaac_env.isaac_env import IsaacEnv
        for _ in range(self.stage_num):
            self.simulator_list.append(
                EnvManager(self.cfg.train, ..., env_cls=IsaacEnv)
            )

环境交互

@register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name="env"),
          blocking=False)
def env_interact_step(self, data: DataProto) -> dict:
    """执行一步环境交互

    接收动作,执行 chunk_step,返回新观测和奖励。
    """
    chunk_actions = data.non_tensor_batch["actions"]
    stage_id = data.meta_info["stage_id"]

    # 执行 chunk step
    extracted_obs, chunk_rewards, chunk_terminations, chunk_truncations, infos = \
        self.simulator_list[stage_id].chunk_step(chunk_actions)

    # 封装为 DataProto
    env_batch = create_env_batch_dataproto(
        obs=extracted_obs, rews=chunk_rewards,
        terminations=chunk_terminations, truncations=chunk_truncations,
        infos=infos
    )
    return env_batch

环境重置

@register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name="env"),
          blocking=False)
def reset_envs_to_state_ids(self, data: DataProto):
    """重置环境到指定状态"""
    state_ids_list = list(data.non_tensor_batch["state_ids"])
    task_ids_list = list(data.non_tensor_batch["task_ids"])

    result_list = []
    for stage_id in range(self.stage_num):
        result = self.simulator_list[stage_id].reset_envs_to_state_ids(
            state_ids_list[stage_id * N : (stage_id + 1) * N],
            task_ids_list[stage_id * N : (stage_id + 1) * N],
        )
        result_list.append(result)

    # 合并各阶段的结果
    output_tensor_dict = {}
    for k in images_and_states_list[0].keys():
        output_tensor_dict[k] = torch.cat([d[k] for d in images_and_states_list])
    return DataProto.from_dict(tensors=output_tensor_dict, ...)

辅助函数

def put_tensor_cpu(data_dict):
    """递归将字典中的所有 tensor 移到 CPU"""
    for key, value in data_dict.items():
        if isinstance(value, torch.Tensor):
            data_dict[key] = value.cpu().contiguous()
    return data_dict

def create_env_batch_dataproto(obs, rews, terminations, truncations, infos, meta=None):
    """将环境输出转换为 DataProto 格式"""
    tensor_batch = {
        "full_image": obs["images_and_states"]["full_image"],
        "wrist_image": obs["images_and_states"]["wrist_image"],
        "state": obs["images_and_states"]["state"],
        "rews": rews,
        "terminations": terminations,
        "truncations": truncations,
    }
    non_tensor_batch = {"task_descriptions": obs["task_descriptions"]}
    return DataProto.from_dict(tensors=tensor_batch, non_tensors=non_tensor_batch)

核心类/函数列表

名称 类型 说明
EnvWorker 类 Ray 环境工作器
init_worker 方法 创建环境管理器
init_simulator 方法 启动仿真器子进程
env_interact_step 方法 执行环境交互
reset_envs_to_state_ids 方法 重置环境状态
get_all_state_ids 方法 获取所有可用状态 ID
finish_rollout 方法 完成一轮 rollout(保存视频等)
put_tensor_cpu 函数 tensor 移到 CPU
create_env_batch_dataproto 函数 环境输出 -> DataProto

与其他模块的关系

  • 使用 EnvManager(env_manager.py)管理仿真器子进程
  • 使用 LiberoEnv 或 IsaacEnv 作为底层环境
  • 被 EnvLoop(env_loop.py)通过 RayWorkerGroup 远程调用
  • 使用 verl 的 @register 装饰器暴露 RPC 接口

小结

EnvWorker 是环境侧的 Ray Actor,负责在 Ray 集群中管理仿真环境。通过 @register 装饰器,它的方法可以被训练器远程调用。每个 Worker 管理多个流水线阶段的环境实例,支持异步交互。