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

Stable-Baselines3自定义环境与Policy的_predict方法类型错误

问题

将自定义环境和Policy集成到Stable-Baselines3(SB3)时,手动调用_predict传入标准Python dict观测能正常运行,但用PPO训练时,传入的是gymnasium.spaces.dict.Dict类型对象而非实际观测,报错:

AttributeError: 'Box' object has no attribute flatten

已用dict表示环境观测、用gym.spaces定义观测与动作空间,需解决该问题,明确是否需要使用SB3特征提取器或遗漏配置。

原因分析
  1. 策略基类使用错误:直接继承BasePolicy未实现SB3内部要求的观测转换逻辑,SB3训练时会传递批量张量观测(而非单个numpy dict),且Dict空间的观测会被封装为张量字典,而非原始numpy数组。
  2. PPO初始化错误:原代码中PPO(policy=custom_policy, env=env, verbose=1).model.learn调用方式错误,应直接实例化PPO对象后调用learn方法,而非访问内部model属性。
  3. forward方法未适配批量张量:自定义策略的forward方法直接处理numpy数组,未适配SB3传递的批量张量输入,也未正确解析Dict类型的张量观测。
解决方案

1. 正确继承SB3策略基类

改用ActorCriticPolicy作为基类(PPO依赖Actor-Critic架构),确保SB3能自动处理观测转换、批量输入等逻辑。

2. 适配Dict类型的张量观测

使用SB3内置的CombinedExtractor自动处理Dict观测空间,将多个Box特征扁平拼接,无需手动处理flatten操作。

3. 修正PPO训练代码

按SB3标准流程初始化PPO并调用训练方法。

修改后的完整代码

import numpy as np
import gymnasium as gym
import torch
import torch.nn as nn
import torch.nn.functional as F
from stable_baselines3 import PPO
from stable_baselines3.common.evaluation import evaluate_policy
from stable_baselines3.common.policies import ActorCriticPolicy
from stable_baselines3.common.torch_layers import CombinedExtractor
from stable_baselines3.common.utils import get_schedule_fn


class CustomEnv(gym.Env):
    def __init__(self, nRows, nCols):
        super().__init__()
        self.nRows = nRows
        self.nCols = nCols
        self.iter = 0
        self.done = False
        self.truncated = False

        self.action_space = gym.spaces.Discrete(self.nRows * self.nCols)
        self.observation_space = gym.spaces.Dict({
            'layout': gym.spaces.Box(low=0, high=255, shape=(self.nRows, self.nCols), dtype=np.uint8),
            'mask': gym.spaces.Box(low=0, high=255, shape=(self.nRows * self.nCols,), dtype=np.uint8)
        })
        self.observation = {'layout': np.zeros((self.nRows, self.nCols), dtype=np.uint8),
                            'mask': np.zeros(self.nRows * self.nCols, dtype=np.uint8)}

    def step(self, action):
        self.iter += 1
        reward = 0
        layout = self.observation["layout"].flatten()
        mask = self.observation["mask"]

        if layout[action] == 0:
            layout[action] = 1
            mask[action] = 1
            reward = 1
        else:
            reward = -1

        self.observation = {'layout': np.reshape(layout, (self.nRows, self.nCols)),
                            'mask': mask}

        if self.iter > self.nRows * self.nCols:
            self.done = True
            self.truncated = True

        return self.observation, reward, self.done, self.truncated, {}

    def reset(self, seed=None, options=None):
        super().reset(seed=seed)
        self.iter = 0
        self.done = False
        self.truncated = False
        self.observation = {'layout': np.zeros((self.nRows, self.nCols), dtype=np.uint8),
                            'mask': np.zeros(self.nRows * self.nCols, dtype=np.uint8)}
        return self.observation, {}

    def render(self):
        pass

    def close(self):
        pass


class CustomPolicy(ActorCriticPolicy):
    def __init__(self, observation_space, action_space, lr_schedule, **kwargs):
        # 用CombinedExtractor自动处理Dict观测空间的特征提取
        super().__init__(
            observation_space,
            action_space,
            lr_schedule,
            features_extractor_class=CombinedExtractor,
            **kwargs
        )

    def _build_mlp_extractor(self) -> None:
        # 自定义MLP网络结构,输入为特征提取后的扁平张量
        feature_dim = self.features_extractor.features_dim
        self.mlp_extractor = nn.Sequential(
            nn.Linear(feature_dim, 5),
            nn.ReLU(),
            nn.Linear(5, self.action_space.n)
        )

    def forward(self, obs, deterministic: bool = False):
        # 提取观测特征
        features = self.extract_features(obs)
        # 生成动作logits
        action_logits = self.mlp_extractor(features)
        # 计算值函数(PPO必须的输出)
        value = self.value_net(features)
        # 构建动作分布并采样
        action_dist = self._get_action_dist_from_logits(action_logits)
        actions = action_dist.get_actions(deterministic=deterministic)
        log_prob = action_dist.log_prob(actions)
        return actions, value, log_prob

    def _predict(self, obs, deterministic: bool = False):
        # 简化预测逻辑,复用特征提取和动作分布逻辑
        action_dist = self._get_action_dist_from_logits(self.mlp_extractor(self.extract_features(obs)))
        return action_dist.get_actions(deterministic=deterministic)


# 创建环境实例
env = CustomEnv(3, 3)

# 手动测试策略(可选)
lr_schedule = get_schedule_fn(3e-4)
custom_policy = CustomPolicy(env.observation_space, env.action_space, lr_schedule)
observation, _ = env.reset()
# 转换观测为SB3内部使用的张量格式
obs_tensor = custom_policy.obs_to_tensor(observation)[0]
action = custom_policy._predict(obs_tensor, deterministic=False)
print("手动测试动作:", action)

# 训练PPO模型
model = PPO(policy=CustomPolicy, env=env, verbose=1, learning_rate=3e-4)
model.learn(total_timesteps=1000)

# 评估模型
mean_reward, std_reward = evaluate_policy(model, env, n_eval_episodes=10)
print(f"Mean reward: {mean_reward} +/- {std_reward}")

关键修改说明

  • 改用ActorCriticPolicy作为基类,完全适配PPO的Actor-Critic架构要求。
  • 引入CombinedExtractor自动处理Dict观测空间,无需手动解析每个Box特征的扁平操作。
  • 修正forward和_predict方法,适配SB3传递的批量张量输入逻辑。
  • 修复PPO初始化代码,按标准流程实例化并调用训练方法。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 06:00:56