OpenAI Gym环境15x15 numpy数组观测空间正确定义方法
问题背景
搭建自定义OpenAI Gym环境时,观测值为15x15网格:网格初始值全为0,运行过程中元素取值范围为0~255;动作空间共225个离散动作,每个动作对应网格上一个位置。
初始编写的__init__方法空间定义代码如下:
self.action_space = Discrete(225) self.observation_shape = Box(low=-1000,high=10000,shape=(15,15,),dtype=np.uint8)
报错现象
运行Stable Baselines 3的DQN相关代码时抛出异常:
import stable_baselines3 from stable_baselines3 import DQN model = DQN("MultiInputPolicy", env, verbose=1) model.learn(total_timesteps=10000, log_interval=4)
报错核心信息:
NotImplementedError: [[0 0 0 0 0 0 0 0 0 0 0 0 0 0 0] [0 0 0 0 0 0 0 0 0 0 0 0 0 0 0] ...(省略全0网格行) [0 0 0 0 0 0 0 0 0 0 0 0 0 0 0]] observation space is not supported
报错触发逻辑来自SB3的观测空间校验代码:
if isinstance(observation_space, spaces.Box): return observation_space.shape elif isinstance(observation_space, spaces.Discrete): # Observation is an int return (1,) elif isinstance(observation_space, spaces.MultiDiscrete): # Number of discrete features return (int(len(observation_space.nvec)),) elif isinstance(observation_space, spaces.MultiBinary): # Number of binary features return (int(observation_space.n),) elif isinstance(observation_space, spaces.Dict): return {key: get_obs_shape(subspace) for (key, subspace) in observation_space.spaces.items()} else: raise NotImplementedError(f"{observation_space} observation space is not supported")
错误原因
一共存在3个问题:
- 核心属性名写错:OpenAI Gym环境强制要求观测空间的属性名必须为
observation_space,代码里写成了observation_shape,导致SB3无法读取到定义的Box空间,实际拿到的是环境返回的numpy数组观测值,因此触发类型不支持的报错,从报错信息打印出全0数组也能印证这一点。 - Box空间参数不匹配:指定
dtype=np.uint8时,数据取值范围只能是0~255,设置low=-1000、high=10000超出数据类型范围,会出现溢出问题。 - Policy类型选错:
MultiInputPolicy仅适用于Dict类型的多输入观测空间,单一数组观测不需要用这个策略。
正确实现方法
- 修正环境的空间定义,注意属性名和参数匹配:
import gym from gym.spaces import Discrete, Box import numpy as np # 环境__init__方法内的定义 self.action_space = Discrete(225) # 注意属性名必须是observation_space,low/high匹配uint8的取值范围 # 如果用多层感知机处理数据,shape直接设为(15,15)即可;如果用卷积神经网络处理,补充通道维设为(15,15,1) self.observation_space = Box(low=0, high=255, shape=(15,15), dtype=np.uint8)
- 选择匹配的Policy类型初始化模型,不要用
MultiInputPolicy:
from stable_baselines3 import DQN # 单一Box观测,用MlpPolicy即可;如果shape设为带通道维的图像格式,替换为CnnPolicy model = DQN("MlpPolicy", env, verbose=1) model.learn(total_timesteps=10000, log_interval=4)
补充说明:Keras-rl2中使用的(1,15,15)是通道在前的维度格式,Stable Baselines 3基于PyTorch实现,默认使用通道在后的格式(高、宽、通道数),不要直接混用维度写法。
内容的提问来源于stack exchange,提问作者Mercury
相关产品推荐
相关产品推荐

