使用PPO训练CarRacing-v2时model.learn()报ValueError错误求助
问题:CarRacing-v2 + PPO训练时
model.learn()触发ValueError 问题背景
尝试使用OpenAI Gym的CarRacing-v2环境,通过PPO算法训练车辆,代码如下:
import os import gym from stable_baselines3 import PPO from stable_baselines3.common.vec_env import DummyVecEnv from stable_baselines3.common.evaluation import evaluate_policy environment_name = 'CarRacing-v2' env = gym.make(environment_name, render_mode='human') env.reset() env.close() env = gym.make(environment_name, render_mode='human') episodes = 5 for episode in range(1, episodes+1): observation, info = env.reset() terminated = False truncated = False score = 0 while not (terminated or truncated): # env.render() action = env.action_space.sample() observation, reward, terminated, truncated, info = env.step(action) score += reward print(f'Episode: {episode} Score: {score}') env.close() env = gym.make(environment_name) env = DummyVecEnv([lambda: env]) log_path = os.path.join('Training', 'Logs') model = PPO('CnnPolicy', env, verbose=1, tensorboard_log=log_path) model.learn(total_timesteps=200000)
报错信息
运行最后一行model.learn(total_timesteps=200000)时,触发如下错误:
ValueError Traceback (most recent call last) Cell In[19], line 1 ----> 1 model.learn(total_timesteps=200000, reset_num_timesteps=False) File ~\AppData\Local\Programs\Python\Python311\Lib\site-packages\stable_baselines3\ppo\ppo.py:299, in PPO.learn(self, total_timesteps, callback, log_interval, eval_env, eval_freq, n_eval_episodes, tb_log_name, eval_log_path, reset_num_timesteps) 286 def learn( 287 self, 288 total_timesteps: int, (...) 296 reset_num_timesteps: bool = True, 297 ) -> "PPO": --> 299 return super(PPO, self).learn( 300 total_timesteps=total_timesteps, 301 callback=callback, 302 log_interval=log_interval, 303 eval_env=eval_env, 304 eval_freq=eval_freq, 305 n_eval_episodes=n_eval_episodes, 306 tb_log_name=tb_log_name, 307 eval_log_path=eval_log_path, 308 reset_num_timesteps=reset_num_timesteps, 309 ) File ~\AppData\Local\Programs\Python\Python311\Lib\site-packages\stable_baselines3\common\on_policy_algorithm.py:242, in OnPolicyAlgorithm.learn(self, total_timesteps, callback, log_interval, eval_env, eval_freq, n_eval_episodes, tb_log_name, eval_log_path, reset_num_timesteps) 228 def learn( 229 self, 230 total_timesteps: int, (...) 238 reset_num_timesteps: bool = True, 239 ) -> "OnPolicyAlgorithm": 240 iteration = 0 --> 242 total_timesteps, callback = self._setup_learn( 243 total_timesteps, eval_env, callback, eval_freq, n_eval_episodes, eval_log_path, reset_num_timesteps, tb_log_name 244 ) 246 callback.on_training_start(locals(), globals()) 248 while self.num_timesteps < total_timesteps: File ~\AppData\Local\Programs\Python\Python311\Lib\site-packages\stable_baselines3\common\base_class.py:429, in BaseAlgorithm._setup_learn(self, total_timesteps, eval_env, callback, eval_freq, n_eval_episodes, log_path, reset_num_timesteps, tb_log_name) 427 # Avoid resetting the environment when calling ``.learn()`` consecutive times 428 if reset_num_timesteps or self._last_obs is None: --> 429 self._last_obs = self.env.reset() # pytype: disable=annotation-type-mismatch 430 self._last_episode_starts = np.ones((self.env.num_envs,), dtype=bool) 431 # Retrieve unnormalized observation for saving into the buffer File ~\AppData\Local\Programs\Python\Python311\Lib\site-packages\stable_baselines3\common\vec_env\vec_transpose.py:110, in VecTransposeImage.reset(self) 106 def reset(self) -> Union[np.ndarray, Dict]: 107 """ 108 Reset all environments 109 """ --> 110 return self.transpose_observations(self.venv.reset()) File ~\AppData\Local\Programs\Python\Python311\Lib\site-packages\stable_baselines3\common\vec_env\dummy_vec_env.py:62, in DummyVecEnv.reset(self) 60 for env_idx in range(self.num_envs): 61 obs = self.envs[env_idx].reset() --> 62 self._save_obs(env_idx, obs) 63 return self._obs_from_buf() File ~\AppData\Local\Programs\Python\Python311\Lib\site-packages\stable_baselines3\common\vec_env\dummy_vec_env.py:92, in DummyVecEnv._save_obs(self, env_idx, obs) 90 for key in self.keys: 91 if key is None: --> 92 self.buf_obs[key][env_idx] = obs 93 else: 94 self.buf_obs[key][env_idx] = obs[key] ValueError: setting an array element with a sequence. The requested array has an inhomogeneous shape after 1 dimensions. The detected shape was (2,) + inhomogeneous part.
使用版本:stable-baselines3 v1.7.0,gym v0.26.0,Python为最新版本。
解决方案
错误原因
Gym 0.26+版本修改了API:env.reset()返回**(observation, info)**元组;而stable-baselines3 v1.7.0仅兼容Gym旧API(reset()只返回observation),DummyVecEnv无法处理新API的返回格式,导致存储观测时出错。
方法一:升级stable-baselines3到兼容版本
stable-baselines3 v1.8.0及以上版本已适配Gym 0.26+的新API,执行以下命令升级:
pip install --upgrade stable-baselines3
升级后无需修改原有代码,直接运行即可。
方法二:不升级,用兼容包装器适配旧API
如果不想升级SB3,可使用Gym自带的CompatibilityV0包装器,让环境reset()仅返回observation,适配SB3 v1.7.0的要求:
修改训练环境初始化部分的代码:
env = gym.make(environment_name) # 添加兼容包装器,将reset()返回值转为仅observation env = gym.wrappers.CompatibilityV0(env, render_mode=None) env = DummyVecEnv([lambda: env])
额外优化
原代码中重复创建了带render_mode='human'的环境,训练阶段不需要渲染,直接创建无渲染模式的环境即可,减少资源占用。
内容的提问来源于stack exchange,提问作者Chirag Sharma
相关产品推荐
相关产品推荐

