Stable Baselines3结合Dict观测空间自定义环境创建PPO模型遇TypeError
问题:Stable Baselines3 结合Dict观测空间与MultiDiscrete动作空间创建模型时触发TypeError
尝试使用Stable Baselines3的PPO/A2C算法结合自定义环境训练模型,环境已通过check_env检查,但创建模型时触发TypeError。环境采用Dict观测空间,搭配MultiInputPolicy,测试后确认问题源于MultiDiscrete动作空间的错误定义。
错误信息
TypeError Traceback (most recent call last) Cell In[7], line 4 1 env = DominoTrainEnv(6) 3 # Initialize the PPO agent ----> 4 model = PPO("MultiInputPolicy", env) 6 # Train the agent 7 model.learn(total_timesteps=10000) File ~\AppData\Roaming\Python\Python311\site-packages\stable_baselines3\ppo\ppo.py:164, in PPO.__init__(self, policy, env, learning_rate, n_steps, batch_size, n_epochs, gamma, gae_lambda, clip_range, clip_range_vf, normalize_advantage, ent_coef, vf_coef, max_grad_norm, use_sde, sde_sample_freq, target_kl, stats_window_size, tensorboard_log, policy_kwargs, verbose, seed, device, _init_setup_model) 161 self.target_kl = target_kl 163 if _init_setup_model: --> 164 self._setup_model() File ~\AppData\Roaming\Python\Python311\site-packages\stable_baselines3\ppo\ppo.py:167, in PPO._setup_model(self) 166 def _setup_model(self) -> None: --> 167 super()._setup_model() 169 # Initialize schedules for policy/value clipping 170 self.clip_range = get_schedule_fn(self.clip_range) File ~\AppData\Roaming\Python\Python311\site-packages\stable_baselines3\common\on_policy_algorithm.py:123, in OnPolicyAlgorithm._setup_model(self) 113 self.rollout_buffer = buffer_cls( 114 self.n_steps, 115 self.observation_space, (...) 120 n_envs=self.n_envs, 121 ) 122 # pytype:disable=not-instantiable --> 123 self.policy = self.policy_class( # type: ignore[assignment] 124 self.observation_space, self.action_space, self.lr_schedule, use_sde=self.use_sde, **self.policy_kwargs 125 ) 126 # pytype:enable=not-instantiable 127 self.policy = self.policy.to(self.device) File ~\AppData\Roaming\Python\Python311\site-packages\stable_baselines3\common\policies.py:853, in MultiInputActorCriticPolicy.__init__(self, observation_space, action_space, lr_schedule, net_arch, activation_fn, ortho_init, use_sde, log_std_init, full_std, use_expln, squash_output, features_extractor_class, features_extractor_kwargs, share_features_extractor, normalize_images, optimizer_class, optimizer_kwargs) 833 def __init__( 834 self, 835 observation_space: spaces.Dict, (...) 851 optimizer_kwargs: Optional[Dict[str, Any]] = None, 852 ): --> 853 super().__init__( 854 observation_space, 855 action_space, 856 lr_schedule, 857 net_arch, 858 activation_fn, 859 ortho_init, 860 use_sde, 861 log_std_init, 862 full_std, 863 use_expln, 864 squash_output, 865 features_extractor_class, 866 features_extractor_kwargs, 867 share_features_extractor, 868 normalize_images, 869 optimizer_class, 870 optimizer_kwargs, 871 ) File ~\AppData\Roaming\Python\Python311\site-packages\stable_baselines3\common\policies.py:507, in ActorCriticPolicy.__init__(self, observation_space, action_space, lr_schedule, net_arch, activation_fn, ortho_init, use_sde, log_std_init, full_std, use_expln, squash_output, features_extractor_class, features_extractor_kwargs, share_features_extractor, normalize_images, optimizer_class, optimizer_kwargs) 504 # Action distribution 505 self.action_dist = make_proba_distribution(action_space, use_sde=use_sde, dist_kwargs=dist_kwargs) --> 507 self._build(lr_schedule) File ~\AppData\Roaming\Python\Python311\site-packages\stable_baselines3\common\policies.py:577, in ActorCriticPolicy._build(self, lr_schedule) 573 self.action_net, self.log_std = self.action_dist.proba_distribution_net( 574 latent_dim=latent_dim_pi, latent_sde_dim=latent_dim_pi, log_std_init=self.log_std_init 575 ) 576 elif isinstance(self.action_dist, (CategoricalDistribution, MultiCategoricalDistribution, BernoulliDistribution)): --> 577 self.action_net = self.action_dist.proba_distribution_net(latent_dim=latent_dim_pi) 578 else: 579 raise NotImplementedError(f"Unsupported distribution '{self.action_dist}'.") File ~\AppData\Roaming\Python\Python311\site-packages\stable_baselines3\common\distributions.py:336, in MultiCategoricalDistribution.proba_distribution_net(self, latent_dim) 325 def proba_distribution_net(self, latent_dim: int) -> nn.Module: 326 """ 327 Create the layer that represents the distribution: 328 it will be the logits (flattened) of the MultiCategorical distribution. (...) 333 :return: 334 """ --> 336 action_logits = nn.Linear(latent_dim, sum(self.action_dims)) 337 return action_logits File ~\AppData\Roaming\Python\Python311\site-packages\torch\nn\modules\linear.py:96, in Linear.__init__(self, in_features, out_features, bias, device, dtype) 94 self.in_features = in_features 95 self.out_features = out_features ---> 96 self.weight = Parameter(torch.empty((out_features, in_features), **factory_kwargs)) 97 if bias: 98 self.bias = Parameter(torch.empty(out_features, **factory_kwargs)) TypeError: empty() received an invalid combination of arguments - got (tuple, dtype=NoneType, device=NoneType), but expected one of: * (tuple of ints size, *, tuple of names names, torch.memory_format memory_format, torch.dtype dtype, torch.layout layout, torch.device device, bool pin_memory, bool requires_grad) * (tuple of ints size, *, torch.memory_format memory_format, Tensor out, torch.dtype dtype, torch.layout layout, torch.device device, bool pin_memory, bool requires_grad)
最小可复现代码
import gymnasium as gym from gymnasium import Env from gymnasium.spaces import Discrete, MultiDiscrete, Box, Dict import numpy as np from stable_baselines3 import PPO from stable_baselines3.common.env_checker import check_env class TestEnv(Env): def __init__(self,numPlayers:int): # Actions we can take, 13,13 for possible domino sides, [9,13] for possible domino placements self.action_space = MultiDiscrete(np.array([[13, 13], [9, 13]])) # observation space obsv = { "hand": Box(high=np.array([[13, 13]*79]), dtype=np.int8,low=np.array([[-1, -1]*79])) } self.observation_space = Dict(obsv) self.state = self.getState() @staticmethod def __padArray(array,len:int): return np.pad(array,((0,len),(0,0)),mode='constant',constant_values=-1) def getState(self): #array values are placeholders handarray = np.array([(11,11),(12,11)], dtype=np.int8) hand_padding = TestEnv.__padArray(handarray, 79-len(handarray)) state = { "hand": hand_padding.ravel().reshape((1,158)), } return state def step(self, action): self.state = self.getState() # Return step information reward = 1 done = False info = {} return self.state, reward, done,False, info def render(self): # Implement viz pass def reset(self, seed=None): self.state = self.getState() return self.state, {} env = TestEnv(6) print(check_env(env,skip_render_check=True)) model = PPO("MultiInputPolicy", env)
解决方案
1. 修正MultiDiscrete动作空间定义
MultiDiscrete要求传入一维数组,每个元素对应一个离散动作维度的可选值数量上限。原代码的二维数组会导致内部计算错误,触发TypeError。
修正后的动作空间定义:
self.action_space = MultiDiscrete(np.array([13, 13, 9, 13]))
2. 修正观测空间的Box定义
Box的high和low参数需要是一维数组,同时返回的观测数据不能额外增加维度,要与空间定义匹配。
修正后的观测空间和状态返回:
# 修正Box的high/low为一维数组 obsv = { "hand": Box(high=np.array([13, 13]*79), dtype=np.int8, low=np.array([-1, -1]*79)) } # 修正getState的返回值为一维数组 state = { "hand": hand_padding.ravel(), }
完整修正代码
import gymnasium as gym from gymnasium import Env from gymnasium.spaces import Discrete, MultiDiscrete, Box, Dict import numpy as np from stable_baselines3 import PPO from stable_baselines3.common.env_checker import check_env class TestEnv(Env): def __init__(self,numPlayers:int): # 修正MultiDiscrete参数为一维数组 self.action_space = MultiDiscrete(np.array([13, 13, 9, 13])) # 修正Box的high/low为一维数组 obsv = { "hand": Box(high=np.array([13, 13]*79), dtype=np.int8, low=np.array([-1, -1]*79)) } self.observation_space = Dict(obsv) self.state = self.getState() @staticmethod def __padArray(array,len:int): return np.pad(array,((0,len),(0,0)),mode='constant',constant_values=-1) def getState(self): handarray = np.array([(11,11),(12,11)], dtype=np.int8) hand_padding = TestEnv.__padArray(handarray, 79-len(handarray)) # 修正返回的hand为一维数组 state = { "hand": hand_padding.ravel(), } return state def step(self, action): self.state = self.getState() reward = 1 done = False info = {} return self.state, reward, done, False, info def render(self): pass def reset(self, seed=None): self.state = self.getState() return self.state, {} env = TestEnv(6) print(check_env(env,skip_render_check=True)) model = PPO("MultiInputPolicy", env)
错误原因说明
MultiDiscrete的构造参数必须是一维数组,原二维数组会导致sum(self.action_dims)计算出错误值,进而在创建PyTorch线性层时传入非法参数,触发TypeError。- 观测空间的
Box要求high和low为一维数组,原二维数组虽通过check_env,但会导致后续模型处理观测数据时出现维度不匹配问题。
内容的提问来源于stack exchange,提问作者lilmrmagoo
相关产品推荐
相关产品推荐

