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

Stable Baselines3 PPO模型训练时因dummy_vec_env.py报错崩溃求助

问题描述

尝试在CartPole-v1环境中训练PPO模型,代码如下:

import gym
from stable_baselines3 import PPO
from stable_baselines3.common.vec_env import DummyVecEnv, VecNormalize
from stable_baselines3.common.env_util import make_vec_env
from stable_baselines3.common.evaluation import evaluate_policy

env_id = "CartPole-v1"
#Making the environment
envs = make_vec_env(env_id, n_envs= 4)
envs = VecNormalize(envs)

#Training the model
model = PPO(policy="MlpPolicy", env=envs, verbose=1)
model.learn(1000)
model.save("CartPole-v1-model")
envs.save("CartPole-v1-env")

运行时触发dummy_vec_env.py相关错误,本地调试发现obs变量为元组;相同代码在HuggingFace的Google Colab笔记本可正常运行。本地安装的是仅支持CPU的PyTorch,怀疑与此有关,但查看dummy_vec_env.py及其父类base_vec_env.py未导入PyTorch,暂时无法确定根因。

解决方案
  • 适配Gym新版本的观测格式
    Gym 0.26+版本中,step()和reset()方法的返回值改为元组((obs, info)或(obs, reward, terminated, truncated, info)),而旧版本仅返回观测/观测+奖励等单一值。Stable-Baselines3旧版本未适配这种格式,会导致DummyVecEnv处理观测时出错。
    解决办法二选一:

    1. 降级Gym到兼容版本:
      pip install gym==0.25.2
      
    2. 添加环境包装器提取观测部分:
      import gym
      from stable_baselines3 import PPO
      from stable_baselines3.common.vec_env import VecNormalize
      from stable_baselines3.common.env_util import make_vec_env
      
      class ObsWrapper(gym.Wrapper):
          def step(self, action):
              obs, reward, terminated, truncated, info = self.env.step(action)
              return obs, reward, terminated or truncated, info
          def reset(self, **kwargs):
              obs, info = self.env.reset(**kwargs)
              return obs
      
      env_id = "CartPole-v1"
      envs = make_vec_env(env_id, n_envs=4, wrapper_class=ObsWrapper)
      envs = VecNormalize(envs)
      
      model = PPO(policy="MlpPolicy", env=envs, verbose=1)
      model.learn(1000)
      model.save("CartPole-v1-model")
      envs.save("CartPole-v1-env")
      
  • 更新Stable-Baselines3到最新版本
    新版本的Stable-Baselines3已经适配了Gym 0.26+的API格式,执行以下命令更新:

    pip install --upgrade stable-baselines3
    
  • 验证PyTorch CPU版本的完整性
    虽然环境代码未直接导入PyTorch,但PPO模型依赖PyTorch运行,若安装不完整可能间接引发异常。执行以下代码验证:

    import torch
    print(torch.cuda.is_available())  # CPU版应返回False
    print(torch.__version__)
    

    若输出异常,重新安装CPU版PyTorch:

    pip3 install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 05:05:28