加载PPO模型触发AssertionError:PyBullet+Gym强化学习问题
问题:加载Stable Baselines3 PPO模型时触发AssertionError
基于PyBullet和Gym框架实现强化学习机器人,用PPO训练AntBulletEnv-v0环境的模型并保存为ZIP文件,但加载模型运行时触发AssertionError,报错指向get_schedule_fn函数中断言value_schedule必须是可调用对象。
代码示例
import gym import pybullet, pybullet_envs import torch as th from stable_baselines3 import PPO from stable_baselines3.common.evaluation import evaluate_policy # 创建环境 env = gym.make('AntBulletEnv-v0') env.render(mode="human") policy_kwargs = dict(activation_fn=th.nn.LeakyReLU, net_arch=[512, 512]) # 初始化并删除模型(模拟加载场景) model = PPO('MlpPolicy', env, learning_rate=0.0003, policy_kwargs=policy_kwargs, verbose=1) del model # 加载训练好的模型(此处触发错误) model = PPO.load("ppo_Ant_saved_model_7.zip") # 运行模型 obs = env.reset() for i in range(100): dones = False game_score = 0 steps = 0 while not dones: action, _states = model.predict(obs, deterministic=True) obs, rewards, dones, info = env.step(action) game_score += rewards steps += 1 env.render() print(f"game {i} steps {steps} game score {game_score:.3f}") obs = env.reset()
报错信息
pybullet build time: Oct 16 2022 21:41:54 Using cpu device Wrapping the env with a `Monitor` wrapper Wrapping the env in a DummyVecEnv. C:\Python\Python395\lib\site-packages\stable_baselines3\common\save_util.py:1 using `custom_objects` argument to replace this object. warnings.warn( C:\Python\Python395\lib\site-packages\stable_baselines3\common\save_util.py:1 model._setup_model() File "C:\Python\Python395\lib\site-packages\stable_baselines3\ppo\ppo.py", line 173, in _setup_model self.clip_range = get_schedule_fn(self.clip_range) File "C:\Python\Python395\lib\site-packages\stable_baselines3\common\utils.py", line 91, in get_schedule_fn assert callable(value_schedule) AssertionError
解决方案
原因
报错是因为加载模型时,clip_range参数被解析为固定数值而非可调用的调度函数,违反了get_schedule_fn的断言要求;此外,加载模型时未传入匹配的环境也可能导致参数初始化异常。
修复方案
方案1:加载时传入环境(推荐)
加载模型时传入训练时使用的环境,让模型自动匹配并初始化所有参数:
# 创建环境后直接传入load方法 env = gym.make('AntBulletEnv-v0') model = PPO.load("ppo_Ant_saved_model_7.zip", env=env)
方案2:用custom_objects覆盖clip_range
手动将clip_range转换为可调用的调度函数,通过custom_objects参数传入:
from stable_baselines3.common.utils import get_schedule_fn model = PPO.load("ppo_Ant_saved_model_7.zip", custom_objects={"clip_range": get_schedule_fn(0.2)})
注:0.2是PPO默认的clip_range值,若训练时修改过该参数,请替换为对应数值。
额外检查
确保训练和加载模型时使用的Stable Baselines3版本一致,版本差异可能导致保存的参数结构不兼容。
内容的提问来源于stack exchange,提问作者Topaz Blue
相关产品推荐
相关产品推荐

