自定义Gym环境reset()返回Tuple触发AssertionError问题排查
问题根源与解决方案
错误直接原因
你在reset()方法里返回了元组(self.obs, {'arg1': 0}),但stable_baselines3的check_env默认遵循旧版Gym(v0.x)的接口规范:reset()方法只能返回单个观测值,不能附带info字典。这就是断言错误的核心原因。
解决步骤
根据你的需求选择以下方案:
方案1:适配旧版Gym接口(兼容check_env默认检查)
修改reset()方法,只返回观测值,去掉info字典:
def reset(self): # 初始化代码 self.obs = self.get_ob() return self.obs # 仅返回观测,不返回info
方案2:保留info返回(适配Gymnasium/Gym v1.x)
如果你的环境是基于Gymnasium(原Gym v1.x)开发的,需要在调用check_env时明确指定兼容Gymnasium接口:
check_env(env, gymnasium=True)
同时确保你的自定义环境继承自gymnasium.Env而非旧版gym.Env。
额外潜在问题提示
你的get_ob()返回的数组形状是(self.camera_width, self.camera_height, 3),但observation_space定义的形状是(self.camera_height, self.camera_width, 3),两者的宽高维度顺序不一致。虽然这不是当前报错的原因,但后续会导致观测值与空间不匹配的错误,建议修正get_ob():
def get_ob(self): # 改成和observation_space一致的形状:(height, width, 3) state = np.zeros((self.camera_height, self.camera_width, 3), np.uint8) return state
内容的提问来源于stack exchange,提问作者thansen0
相关产品推荐
相关产品推荐

