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

Gym 0.21+Stable Baselines3训练自定义环境触发TypeError报错

解决自定义Gym环境训练SB3 PPO时的TypeError问题

问题根源

你使用的是Gym 0.21.0(旧API版本),但自定义环境的reset()和step()方法采用了Gym 0.26+的新API返回格式:

  • reset()返回(observation, info),而旧API仅要求返回observation
  • step()返回(observation, reward, terminated, truncated, info),旧API要求返回(observation, reward, done, info)

这导致Stable Baselines3将reset()返回的tuple误判为observation,后续尝试用字符串索引该tuple时触发TypeError。另外你的_agent_location生成时形状错误((2,2)),与观察空间定义的(2,)不匹配,也会引发后续问题。

修复步骤

1. 修正reset方法的返回格式与变量形状

修改gym_envs/envs/grid_world.py中的reset方法,适配旧API并修正位置变量的形状:

def reset(self, seed=None, options=None):
    # 修正agent和target的形状:从(2,2)改为(2,),匹配观察空间定义
    self._agent_location = np.random.randint(0, self.size, size=2)
    self._target_location = self._agent_location
    while np.array_equal(self._target_location, self._agent_location):
        self._target_location = np.random.randint(0, self.size, size=2)

    observation = self._get_obs()
    info = self._get_info()
    if self.render_mode == "human":
        self._render_frame()
    # 旧API仅返回observation
    return observation

2. 修正step方法的返回格式

修改step方法,适配旧API的返回结构:

def step(self, action):
    direction = self._action_to_direction[action]
    self._agent_location = np.clip(
        self._agent_location + direction, 0, self.size - 1
    )
    terminated = np.array_equal(self._agent_location, self._target_location)
    reward = 1 if terminated else 0
    observation = self._get_obs()
    info = self._get_info()

    if self.render_mode == "human":
        self._render_frame()

    # 旧API返回(observation, reward, done, info),done=terminated(此处无截断逻辑)
    return observation, reward, terminated, info

3. 验证修复

修改后,无论是向量化环境还是单环境都能正常训练:

  • 向量化环境代码保持不变
  • 单环境测试代码修正PPO的导入(小写ppo改为大写PPO):
import gym
from stable_baselines3 import PPO

env = gym.make("gym_envs:gym_envs/GridWorld-v0")

model = PPO("MultiInputPolicy", env, tensorboard_log="./logs/", verbose=1)
model.learn(total_timesteps=1000)

额外说明

如果后续想升级到Gym新API(0.26+),需要:

  • 升级Gym版本至0.26+
  • 确保Stable Baselines3版本支持(SB3 1.7.0已兼容新API)
  • 保留当前reset()和step()的新API返回格式

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 17:21:08