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

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文件。
  • 解决操作:
    1. 保存权重时,单独导出policy的state_dict:
      torch.save(model.policy.state_dict(), "C:/Users/prath/Downloads/TeamProject_v6/python/modelv1/policy_weights.pth")
      
    2. 加载时,先初始化与预训练模型结构一致的模型,再加载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"))
      

4. 验证加载结果

加载完成后,可通过以下方式验证:

  • 打印模型参数,对比加载前后的数值变化;
  • 在环境中执行一步预测,观察输出是否符合预训练模型的预期表现;
  • 查看模型的训练日志,确认后续训练的奖励曲线是否延续预训练的趋势。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 18:42:02