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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 06:22:47