解决gym与Stable-Baselines3代码运行时的数组维度超限ValueError
解决CartPole-v1运行时的ValueError数组维度错误
你遇到的错误源于未训练的模型返回的动作格式与VecEnv的step方法不兼容,或是环境包装后的维度处理问题,以下是几种可行的解决办法:
- 方法一:先训练模型再运行
取消注释model.learn(total_timesteps=10000,)这一行,让PPO模型完成训练后再执行预测逻辑。训练后的模型会输出符合VecEnv要求的动作格式,代码可正常运行:
import gym from stable_baselines3 import PPO env = gym.make("CartPole-v1", render_mode="human") model = PPO("MlpPolicy", env, verbose=1) model.learn(total_timesteps=10000,) # 取消注释训练模型 vec_env = model.get_env() obs = vec_env.reset() print(obs) for i in range(1000): action, _states = model.predict(obs, deterministic=True) obs, reward, done, info = vec_env.step(action) vec_env.render() vec_env.close() # 改为关闭包装后的VecEnv而非原环境
- 方法二:不训练,直接使用随机动作
如果仅需测试环境运行效果,不需要模型预测,可直接通过环境的动作空间采样随机动作,替换掉model.predict部分:
import gym from stable_baselines3 import PPO env = gym.make("CartPole-v1", render_mode="human") model = PPO("MlpPolicy", env, verbose=1) vec_env = model.get_env() obs = vec_env.reset() print(obs) for i in range(1000): action = vec_env.action_space.sample() # 采样随机动作 obs, reward, done, info = vec_env.step(action) vec_env.render() vec_env.close()
- 方法三:手动调整动作格式(不推荐)
若坚持使用未训练的模型预测,可手动将动作转为一维数组,避免维度不匹配:
import gym import numpy as np from stable_baselines3 import PPO env = gym.make("CartPole-v1", render_mode="human") model = PPO("MlpPolicy", env, verbose=1) vec_env = model.get_env() obs = vec_env.reset() print(obs) for i in range(1000): action, _states = model.predict(obs, deterministic=True) action = np.array(action).flatten() # 调整动作维度适配VecEnv obs, reward, done, info = vec_env.step(action) vec_env.render() vec_env.close()
内容的提问来源于stack exchange,提问作者AI ML
相关产品推荐
相关产品推荐

