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

自定义Gym环境添加观测空间报错:输入维度不匹配求助

输入维度不匹配问题解决方案

核心问题拆解

  1. 观测空间定义与实际返回的state维度完全不符:你定义的observation_space是5维Box,但reset返回的是嵌套二维数组,且和你需要的导航观测特征数量不匹配。
  2. 模型输入形状与观测维度不匹配:模型指定input_shape=[2],但实际需要的导航观测特征是7维(目标XY+小车XY+朝向+到目标距离+目标方位),且输入被错误包装成4维结构。
  3. 动作空间定义错误:你的动作是离散的3个选项(左转/直行/右转),却误用了连续动作空间的Box类型。

具体修复代码

1. 修正环境的__init__和reset方法

先将所有导航观测特征扁平化为一维数组,匹配观测空间定义:

def __init__(self):
    # 根据实际场景调整各特征的取值范围
    low = np.array([-10, -10,  # 目标坐标(X,Y)
                    -10, -10,  # 小车坐标(X,Y)
                    -np.pi,    # 小车朝向角度(-π到π)
                    0,         # 到目标的最小距离
                    -np.pi])   # 目标相对方位(-π到π)
    high = np.array([10, 10,
                     10, 10,
                     np.pi,
                     20,  # 到目标的最大距离
                     np.pi])
    # 定义7维的观测空间
    self.observation_space = gym.spaces.Box(low=low, high=high, dtype=np.float32)
    # 动作是离散的3个选项,改用Discrete空间
    self.action_space = gym.spaces.Discrete(3)

def reset(self):
    # 初始化所有导航观测特征,合并为一维数组
    target_pos = np.random.uniform(-5, 5, size=2)
    car_pos = np.random.uniform(-5, 5, size=2)
    heading = np.random.uniform(-np.pi, np.pi)
    dist_to_target = np.linalg.norm(target_pos - car_pos)
    bearing_to_target = np.arctan2(target_pos[1]-car_pos[1], target_pos[0]-car_pos[0]) - heading
    
    # 扁平化所有特征为一维数组
    self.state = np.concatenate([target_pos, car_pos, [heading], [dist_to_target], [bearing_to_target]])
    self.shower_length = 60000
    # 确保返回类型与观测空间一致
    return self.state.astype(np.float32)

2. 修正模型的输入形状

模型输入形状要与观测空间的维度(7维)完全匹配:

def build_model(states, actions):
    model = Sequential()
    # input_shape对应观测特征的总数量(7)
    model.add(Dense(units=24, activation='relu', input_shape=(states,)))
    model.add(Dense(units=24, activation='relu'))
    model.add(Dense(actions, activation='linear'))
    return model

# 调用时传入正确的维度参数
model = build_model(states=7, actions=3)

3. 清除输入冗余维度

如果仍出现多维输入报错,在训练循环中对观测值做维度压缩:

state = env.reset()
# 去掉多余的嵌套维度,确保形状为(7,)
state = np.squeeze(state)

额外注意事项

  • 确保所有观测值的取值范围严格匹配observation_space定义的low/high,避免训练时出现异常值。
  • 改用离散动作空间后,训练算法(如DQN)需适配离散动作输出逻辑。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 10:10:49