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

使用Pickle保存Stable Baselines3训练的A2C模型遇AttributeError求助

问题:Pickle保存Stable Baselines3智能体时触发AttributeError

我正在进行一项机器学习课程项目,需要保存一个包含复杂内容的智能体(agent)对象。尝试使用pickle保存时出现错误:

AttributeError: Can't pickle local object 'constant_fn..func'

代码片段如下:

from finrl.agents.stablebaselines3.models import DRLAgent
import pickle
import os

if os.path.isfile("./filename_pi.obj"):
    print("-FILE FOUND-")
    file_pi = open('filename_pi.obj', 'rb')
    trained_a2c = pickle.load(file_pi)
    file_pi.close()
else:
    print("-FILE NOT FOUND-")
    #A2C
    print("Training A2C model")
    agent = DRLAgent(env=env_train)
    model_a2c = agent.get_model("a2c")
    trained_a2c = agent.train_model(model=model_a2c, tb_log_name="a2c", total_timesteps=50000)
    file_pi = open('filename_pi.obj', 'wb') 
    pickle.dump(trained_a2c, file_pi)
    file_pi.close()

查阅类似问题后,我了解到问题源于非全局对象,但无法修改库中的.get_model和.train_model方法,请问有什么解决办法?是否可不传入trained_a2c,或更换方案?


解决办法

1. 使用Stable Baselines3官方的保存/加载方法(推荐)

Stable Baselines3的模型自带专门的保存与加载接口,完全不需要依赖pickle,这是最适配的方案:

from finrl.agents.stablebaselines3.models import DRLAgent
from stable_baselines3 import A2C
import os

MODEL_PATH = "./a2c_model.zip"

if os.path.isfile(MODEL_PATH):
    print("-FILE FOUND-")
    trained_a2c = A2C.load(MODEL_PATH, env=env_train)
else:
    print("-FILE NOT FOUND-")
    print("Training A2C model")
    agent = DRLAgent(env=env_train)
    model_a2c = agent.get_model("a2c")
    trained_a2c = agent.train_model(model=model_a2c, tb_log_name="a2c", total_timesteps=50000)
    # 用官方方法保存模型
    trained_a2c.save(MODEL_PATH)

官方方法会正确保存模型结构、权重、优化器状态等所有必要信息,从根源避免pickle的序列化限制。

2. 用dill替代pickle(备选)

dill是pickle的扩展库,支持更多类型的对象序列化,包括局部函数。先安装dill:

pip install dill

然后替换代码中的pickle为dill即可:

import dill

# 保存模型
file_pi = open('filename_pi.obj', 'wb') 
dill.dump(trained_a2c, file_pi)
file_pi.close()

# 加载模型
file_pi = open('filename_pi.obj', 'rb')
trained_a2c = dill.load(file_pi)
file_pi.close()

注意:这种方法不如官方方案可靠,后续加载时可能出现环境依赖或版本兼容问题。

3. 仅保存模型权重(进阶)

如果只需要保留模型的预测能力,可以只保存权重参数,后续重新初始化模型后加载权重:

import torch

# 保存模型权重
torch.save(trained_a2c.policy.state_dict(), "./a2c_weights.pth")

# 加载权重:先重建模型结构,再加载权重
agent = DRLAgent(env=env_train)
model_a2c = agent.get_model("a2c")
model_a2c.policy.load_state_dict(torch.load("./a2c_weights.pth"))
trained_a2c = model_a2c

这种方法轻量化,但需要确保重建的模型结构和训练时完全一致。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 18:35:40