PyTorch模型Pickle加载报错:无法获取twoD_predict属性
PyTorch模型Pickle加载属性错误解决方法
错误原因
报错AttributeError: Can't get attribute "twoD_predict" on <module '__main__' from '/load_model.py'>的核心原因是:Pickle序列化保存的是模型对象的引用,而非完整的类代码。加载时,当前运行环境(load_model.py的__main__模块)找不到twoD_predict类的定义,导致无法重建对象。
解决方案
方案1:在加载脚本中导入模型类
如果模型类定义在单独的文件(比如model_def.py),直接在load_model.py中导入该类:
from model_def import twoD_predict import pickle if __name__ == '__main__': with open("model.sav", 'rb') as file: model = pickle.load(file)
如果模型类和训练脚本在同一文件,可通过导入训练脚本的方式引入类:
from train_script import twoD_predict # train_script是你的训练文件名 import pickle if __name__ == '__main__': with open("model.sav", 'rb') as file: model = pickle.load(file)
方案2:改用PyTorch官方推荐的状态字典保存方式
PyTorch官方不建议直接用Pickle保存整个模型对象,更推荐保存模型的state_dict(仅包含参数和缓冲区),这种方式更稳定,不受类定义位置变化的影响:
- 保存模型时:
# 替换原Pickle保存代码 torch.save(obj1.state_dict(), f"{column_names[0]}.pth") - 加载模型时:
import torch from model_def import twoD_predict # 必须导入模型类定义 if __name__ == '__main__': model = twoD_predict() # 先实例化模型 model.load_state_dict(torch.load("model.pth")) # 加载参数 model.eval() # 切换到评估模式(预测前必备)
方案3:临时方案——在加载脚本中复制模型类定义
如果不想拆分文件,可直接在load_model.py中复制twoD_predict类的完整定义(必须和训练时的类代码完全一致):
import torch import torch.nn as nn import torch.nn.functional as F import pickle class twoD_predict(nn.Module): def __init__(self): super().__init__() # 与原始类__init__完全一致的代码 def forward(self, x): # 与原始类forward完全一致的代码 ... def train(self, epochs=100): # 与原始类train完全一致的代码 ... if __name__ == '__main__': with open("model.sav", 'rb') as file: model = pickle.load(file)
内容的提问来源于stack exchange,提问作者ShyaliC51
相关产品推荐
相关产品推荐

