You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何让PyTorch SyncVectorEnv适配CleanRL所需的机器人环境字典观测空间?

解决CleanRL适配字典型观测机器人环境的方法

方法一:自定义环境包装器转换观测格式

直接写一个轻量包装器,把字典型观测拼接成一维张量,同时替换观测空间为具备shape属性的Box类型:

import gym
from gym import spaces
import numpy as np

class DictToFlatObsWrapper(gym.Wrapper):
    def __init__(self, env):
        super().__init__(env)
        # 遍历字典观测空间,统计所有维度
        obs_spaces = []
        for space in env.observation_space.spaces.values():
            if isinstance(space, spaces.Box):
                obs_spaces.append(space.shape)
        # 计算拼接后的总维度
        total_dim = sum(np.prod(shape) for shape in obs_spaces)
        # 重新定义观测空间为Box类型
        self.observation_space = spaces.Box(
            low=-np.inf, high=np.inf, shape=(total_dim,), dtype=np.float32
        )

    def observation(self, obs):
        # 把字典里的每个值扁平化后拼接
        flat_obs = []
        for value in obs.values():
            flat_obs.append(value.flatten())
        return np.concatenate(flat_obs, axis=0)

使用时将机器人环境用该包装器包裹,再传入SyncVectorEnv:

from torchrl.envs import SyncVectorEnv

# 初始化你的机器人环境
def make_env():
    env = YourRobotEnv()
    env = DictToFlatObsWrapper(env)
    return env

vec_env = SyncVectorEnv([make_env for _ in range(num_envs)])

方法二:使用现成的扁平化工具

如果你的环境基于gymnasium(新版gym),可以直接用官方提供的FlattenObservation包装器,自动处理字典或嵌套结构的观测:

from gymnasium.wrappers import FlattenObservation

def make_env():
    env = YourRobotEnv()
    env = FlattenObservation(env)
    return env

该包装器会自动将字典内所有观测值扁平化拼接,同时更新观测空间为带shape属性的Box类型,直接适配CleanRL的要求。

方法三:修改模型适配字典输入(可选)

若不想改变观测格式,可直接调整CleanRL的模型代码,让它接收字典输入并分别处理每个键对应的张量,最后融合特征:

import torch.nn as nn
import torch

class CustomModel(nn.Module):
    def __init__(self, obs_space):
        super().__init__()
        # 为每个字典键定义子网络
        self.subnets = nn.ModuleDict()
        for key, space in obs_space.spaces.items():
            self.subnets[key] = nn.Sequential(
                nn.Linear(np.prod(space.shape), 64),
                nn.ReLU(),
                nn.Linear(64, 64)
            )
        self.fc = nn.Linear(64 * len(obs_space.spaces), 128)

    def forward(self, obs):
        features = []
        for key, subnet in self.subnets.items():
            feat = subnet(obs[key].flatten(start_dim=1))
            features.append(feat)
        fused = torch.cat(features, dim=1)
        return self.fc(fused)

这种方法无需修改环境,但要调整CleanRL的模型定义部分,适合需要保留原始观测结构的场景。

内容的提问来源于stack exchange,提问作者Dirk

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.05 09:45:29