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

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类型的多输入观测空间,单一数组观测不需要用这个策略。
正确实现方法
  1. 修正环境的空间定义,注意属性名和参数匹配:
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)
  1. 选择匹配的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 17:42:43