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特征提取器或遗漏配置。
原因分析
- 策略基类使用错误:直接继承
BasePolicy未实现SB3内部要求的观测转换逻辑,SB3训练时会传递批量张量观测(而非单个numpy dict),且Dict空间的观测会被封装为张量字典,而非原始numpy数组。 - PPO初始化错误:原代码中
PPO(policy=custom_policy, env=env, verbose=1).model.learn调用方式错误,应直接实例化PPO对象后调用learn方法,而非访问内部model属性。 - 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
相关产品推荐
相关产品推荐

