使用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
相关产品推荐
相关产品推荐

