stable_baselines3加载PPO模型报Box对象无shape属性错误
问题描述
在Colab平台基于stable_baselines3训练PPO模型时,使用如下代码完成模型保存:
model.save("model")
模型在Colab平台可正常加载,但在本地环境加载时触发AttributeError报错,加载代码如下:
m = PPO.load("model", env=env)
完整报错栈:
AttributeError Traceback (most recent call last) /tmp/ipykernel_25649/121834194.py in <module> 2 env = e.MinitaurBulletEnv(render=False) 3 env.reset() ----> 4 m2 = PPO.load("model", env=env) 5 for episode in range(1, 6): 6 obs = env.reset() ~/anaconda3/lib/python3.8/site-packages/stable_baselines3/common/base_class.py in load(cls, path, env, device, custom_objects, **kwargs) 668 env = cls._wrap_env(env, data["verbose"]) 669 # Check if given env is valid ---> 670 check_for_correct_spaces(env, data["observation_space"], data["action_space"]) 671 else: 672 # Use stored env, if one exists. If not, continue as is (can be used for predict) ~/anaconda3/lib/python3.8/site-packages/stable_baselines3/common/utils.py in check_for_correct_spaces(env, observation_space, action_space) 217 :param action_space: Action space to check against 218 """ ---> 219 if observation_space != env.observation_space: 220 raise ValueError(f"Observation spaces do not match: {observation_space} != {env.observation_space}") 221 if action_space != env.action_space: ~/anaconda3/lib/python3.8/site-packages/gym/spaces/box.py in __eq__(self, other) 138 139 def __eq__(self, other): ---> 140 return isinstance(other, Box) and (self.shape == other.shape) and np.allclose(self.low, other.low) and np.allclose(self.high, other.high) AttributeError: 'Box' object has no attribute 'shape'
本地使用PyBullet提供的MinitaurBulletEnv环境,初始化代码如下:
import pybullet_envs.bullet.minitaur_gym_env as e import gym env = e.MinitaurBulletEnv(render=False) env.reset()
报错根因
报错核心是本地与Colab训练环境的依赖版本不匹配:
- 高版本gym(>=0.22.0)和旧版pybullet_envs兼容性极差,旧版pybullet_envs的Minitaur环境在实例化时,没有按照高版本gym的Box空间规范初始化
shape属性,导致两个Box空间做相等校验时触发属性缺失错误。 - 模型保存时会序列化训练环境的观测空间、动作空间元信息,加载时会和传入的本地环境做严格校验,两边依赖版本不一致时,不仅会触发属性错误,还可能出现空间维度、范围不匹配导致的推理异常。
解决方案
按优先级推荐以下处理方式:
- 优先对齐全链路依赖版本:在Colab环境中执行以下代码,打印训练环境所有相关依赖的精确版本:
import gym import stable_baselines3 import pybullet import pybullet_envs import numpy as np print(gym.__version__, stable_baselines3.__version__, pybullet.__version__, pybullet_envs.__version__, np.__version__)
本地虚拟环境中安装和Colab完全一致的版本即可,注意多数能正常跑pybullet_envs的环境搭配为gym==0.21.0,高版本gym必然触发各类空间初始化异常。
- 临时调试可手动补全环境空间的缺失属性(仅用于快速验证,不保证推理结果正确性):在加载模型前执行以下代码补全属性:
import numpy as np # 补全观测空间缺失属性 env.observation_space.shape = env.observation_space.low.shape env.observation_space.dtype = np.float32 # 补全动作空间缺失属性 env.action_space.shape = env.action_space.low.shape env.action_space.dtype = np.float32 # 补全后再加载模型 model = PPO.load("model", env=env)
- 调整代码执行顺序:实例化环境后先校验空间属性完整性,再加载模型,加载完成后再执行
env.reset()和推理循环,避免部分旧版本环境在reset前属性未初始化完全的问题。
内容的提问来源于stack exchange,提问作者abdelmoumen
相关产品推荐
相关产品推荐

