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仅要求返回observationstep()返回(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
相关产品推荐
相关产品推荐

