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

Stable-Baselines3中MultiDiscrete与Box空间的自定义策略适配问题

解决Stable-Baselines3中MultiDiscrete动作空间与Box观测空间维度不匹配问题

问题核心

在结合MultiDiscrete动作空间和(5,5)形状的Box观测空间时,默认策略输出维度(25维)与环境观测维度不匹配,尝试自定义策略或特征提取器时出现多种报错(如TypeError: 'CustomPolicy' object is not callable、ValueError: could not broadcast input array from shape (25,) into shape (5,5))。

关键错误点

  1. 环境观测不一致:GridWorldEnv中observation_space定义为(5,5)的Box,但reset方法返回flatten后的25维数组,导致策略接收的观测形状与定义不符。
  2. 自定义策略方式错误:Stable-Baselines3的策略需要继承特定基类(如ActorCriticPolicy),直接定义的CustomPolicy不符合框架要求,无法被PPO调用。
  3. 混淆特征提取器与完整策略:BaseFeaturesExtractor仅用于观测特征预处理,不能作为完整策略使用,错误地将其当作策略类继承。

正确解决方案

步骤1:修正环境的观测一致性

确保环境返回的观测形状与observation_space定义完全一致:

import numpy as np
import gym
import torch.nn as nn
import torch as th
from stable_baselines3 import PPO
from stable_baselines3.common.torch_layers import BaseFeaturesExtractor

class GridWorldEnv(gym.Env):
    def __init__(self):
        self.observation_space = gym.spaces.Box(low=0, high=1, shape=(5, 5), dtype=np.float32)
        self.action_space = gym.spaces.MultiDiscrete([5, 3])  # 5方向 + 3距离

        self.state = np.zeros((5, 5))
        self.state[0, 0] = 1  # 起始位置
        self.goal = (4, 4)  # 目标位置
        self.steps = 0

    def reset(self):
        self.state = np.zeros((5, 5))
        self.state[0, 0] = 1
        self.steps = 0
        # 返回与observation_space一致的(5,5)形状观测
        return self.state.astype(np.float32)

    def step(self, action):
        direction, distance = action
        reward = -1
        done = False

        # 计算移动偏移
        if direction == 0:
            offset = (distance, 0)
        elif direction == 1:
            offset = (-distance, 0)
        elif direction == 2:
            offset = (0, distance)
        elif direction == 3:
            offset = (0, -distance)
        else:
            offset = (0, 0)

        # 更新位置
        current_pos = np.argwhere(self.state == 1)[0]
        new_pos = tuple(np.clip(current_pos + offset, 0, 4))
        self.state[current_pos] = 0
        self.state[new_pos] = 1

        # 检查是否到达目标
        if np.array_equal(new_pos, self.goal):
            reward = 10
            done = True

        self.steps += 1
        if self.steps >= 50:
            done = True

        # 返回(5,5)形状的观测
        return self.state.astype(np.float32), reward, done, {}

步骤2:自定义特征提取器处理观测形状

用BaseFeaturesExtractor将(5,5)的观测flatten为25维,适配MLP策略的输入要求:

class CustomFeaturesExtractor(BaseFeaturesExtractor):
    def __init__(self, observation_space: gym.spaces.Box, features_dim: int = 25):
        super().__init__(observation_space, features_dim)
        # 仅做flatten处理,将(5,5)转为25维
        self.flatten = nn.Flatten()

    def forward(self, observations: th.Tensor) -> th.Tensor:
        # 输入形状:(batch_size, 5, 5),输出形状:(batch_size, 25)
        return self.flatten(observations)

步骤3:结合默认策略与自定义特征提取器

通过policy_kwargs指定特征提取器,使用MlpPolicy即可正常训练:

if __name__ == '__main__':
    env = GridWorldEnv()

    # 配置策略参数,指定自定义特征提取器
    policy_kwargs = dict(
        features_extractor_class=CustomFeaturesExtractor,
        features_extractor_kwargs=dict(features_dim=25),
        # 可选:自定义网络结构
        net_arch=dict(pi=[64, 64], vf=[64, 64])
    )

    # 使用默认MlpPolicy,配合自定义特征提取器
    model = PPO("MlpPolicy", env=env, verbose=1, policy_kwargs=policy_kwargs)
    # 训练1000步
    model.learn(total_timesteps=1000)

    # 测试训练后的模型
    obs = env.reset()
    for _ in range(50):
        action, _states = model.predict(obs)
        obs, rewards, dones, info = env.step(action)
        if dones:
            break

额外说明

如果需要深度自定义策略逻辑,需继承stable_baselines3.common.policies.ActorCriticPolicy并实现_build_mlp_extractor等方法,但大部分场景下,使用自定义特征提取器配合默认策略已足够解决维度匹配问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 00:50:23