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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 20:57:02