Stable Baselines3加载自定义预训练模型失败问题求助
解决Stable Baselines3加载预训练模型的序列化与KeyError问题
问题现象
加载Stable Baselines3预训练模型时,出现cloudpickle反序列化警告,最终触发KeyError: 'policy_class';尝试直接加载.pkl文件无报错,但预训练模型权重未成功加载。
错误日志
C:\Users\prath\miniconda3\envs\rlunitybutler\lib\site-packages\stable_baselines3\common\save_util.py:166: UserWarning: Could not deserialize object policy_class. Consider using `custom_objects` argument to replace this object. Exception: Can't get attribute '_make_function' on <module 'cloudpickle.cloudpickle' from 'C:\\Users\\prath\\miniconda3\\envs\\rlunitybutler\\lib\\site-packages\\cloudpickle\\cloudpickle.py'> warnings.warn( C:\Users\prath\miniconda3\envs\rlunitybutler\lib\site-packages\stable_baselines3\common\save_util.py:166: UserWarning: Could not deserialize object lr_schedule. Consider using `custom_objects` argument to replace this object. Exception: Can't get attribute '_make_function' on <module 'cloudpickle.cloudpickle' from 'C:\\Users\\prath\\miniconda3\\envs\\rlunitybutler\\lib\\site-packages\\cloudpickle\\cloudpickle.py'> warnings.warn( C:\Users\prath\miniconda3\envs\rlunitybutler\lib\site-packages\stable_baselines3\common\save_util.py:166: UserWarning: Could not deserialize object clip_range. Consider using `custom_objects` argument to replace this object. Exception: Can't get attribute '_make_function' on <module 'cloudpickle.cloudpickle' from 'C:\\Users\\prath\\miniconda3\\envs\\rlunitybutler\\lib\\site-packages\\cloudpickle\\cloudpickle.py'> warnings.warn( Wrapping the env in a DummyVecEnv. Traceback (most recent call last): File "C:\Users\prath\Downloads\TeamProject_v6\python\trainagent.py", line 243, in <module> model = A2C.load("C:/Users/prath/Downloads/TeamProject_v6/python/modelv1/Model_v2", env=env, policy=CustomActorCriticPolicy, verbose=1, n_steps=4096, batch_size=128, seed=51, tensorboard_log=f"C:/Users/prath/Downloads/TeamProject_v6/python/ppo", n_epochs=15) File "C:\Users\prath\miniconda3\envs\rlunitybutler\lib\site-packages\stable_baselines3\common\base_class.py", line 708, in load policy=data["policy_class"], KeyError: 'policy_class'
代码上下文
自定义策略类及模型加载代码:
class CustomActorCriticPolicy(ActorCriticPolicy): def __init__( self, observation_space: spaces.Space, action_space: spaces.Space, lr_schedule: Callable[[float], float], *args, **kwargs, ): super().__init__( observation_space, action_space, lr_schedule, # Pass remaining arguments to base class *args, **kwargs, ) # Disable orthogonal initialization self.ortho_init = False def _build_mlp_extractor(self) -> None: self.mlp_extractor = CustomNetwork(self.features_dim) if __name__ == '__main__': env = GoEnv() env = Monitor(env) model = PPO.load("C:/Users/prath/Downloads/TeamProject_v6/python/modelv1/Model_v2", env=env, policy=CustomActorCriticPolicy, verbose=1, n_steps=4096, batch_size=128, seed=51, tensorboard_log=f"C:/Users/prath/Downloads/TeamProject_v6/python/ppo", n_epochs=15) for i in tqdm(range(4096*3*100)): # Perform a training step model.learn(total_timesteps=4096*3, progress_bar=True) model.save("C:/Users/prath/Downloads/TeamProject_v6/python/modelv1/agento"+str(i))
尝试加载.pkl文件的代码:
model = PPO("C:/Users/prath/Downloads/TeamProject_v6/python/modelv1/agento_v1e5.pkl", env, verbose=1, n_steps=4096, batch_size=64, seed=51, tensorboard_log=f"runs/ppo", n_epochs=15) model.policy.load("C:/Users/prath/Downloads/TeamProject_v6/python/modelv1/agento.pkl")
解决方案
1. 统一cloudpickle版本
- 问题根源:不同版本cloudpickle的序列化格式不兼容,
_make_function是旧版本cloudpickle的内部函数,新版本已移除,导致反序列化失败。 - 解决操作:安装与保存模型时一致的cloudpickle版本,例如:
pip install cloudpickle==2.2.0 # 替换为模型保存时的具体版本
2. 用custom_objects显式指定自定义对象
- 问题根源:自定义策略类、动态生成的
lr_schedule和clip_range无法自动反序列化,需要显式告知加载函数替换这些对象。 - 解决操作:加载模型时传入
custom_objects参数,示例代码:from stable_baselines3.common.utils import get_schedule_fn # 定义与训练时一致的学习率调度器和裁剪范围 def custom_lr_schedule(progress_remaining): return progress_remaining * 3e-4 # 替换为你的训练初始学习率 custom_clip_range = 0.2 # 替换为你的训练clip_range值 model = PPO.load( "C:/Users/prath/Downloads/TeamProject_v6/python/modelv1/Model_v2", env=env, custom_objects={ "policy_class": CustomActorCriticPolicy, "lr_schedule": get_schedule_fn(custom_lr_schedule), "clip_range": custom_clip_range }, verbose=1, n_steps=4096, batch_size=128, seed=51, tensorboard_log="C:/Users/prath/Downloads/TeamProject_v6/python/ppo", n_epochs=15 )
3. 正确加载权重文件
- 问题根源:直接用
PPO()构造函数传入.pkl路径是错误用法,model.policy.load()需要加载的是PyTorch的state_dict,而非整个模型的pkl文件。 - 解决操作:
- 保存权重时,单独导出policy的state_dict:
torch.save(model.policy.state_dict(), "C:/Users/prath/Downloads/TeamProject_v6/python/modelv1/policy_weights.pth") - 加载时,先初始化与预训练模型结构一致的模型,再加载state_dict:
import torch # 初始化模型(结构必须与预训练模型完全一致) model = PPO( policy=CustomActorCriticPolicy, env=env, verbose=1, n_steps=4096, batch_size=64, seed=51, tensorboard_log="runs/ppo", n_epochs=15 ) # 加载预训练权重 model.policy.load_state_dict(torch.load("C:/Users/prath/Downloads/TeamProject_v6/python/modelv1/policy_weights.pth"))
- 保存权重时,单独导出policy的state_dict:
4. 验证加载结果
加载完成后,可通过以下方式验证:
- 打印模型参数,对比加载前后的数值变化;
- 在环境中执行一步预测,观察输出是否符合预训练模型的预期表现;
- 查看模型的训练日志,确认后续训练的奖励曲线是否延续预训练的趋势。
内容的提问来源于stack exchange,提问作者hogpip
相关产品推荐
相关产品推荐

