使用pickle加载BERT NER模型失败,报AttributeError错误求助
解决PyTorch BERT NER模型pickle加载报错问题
问题原因
pickle序列化时仅保存类的引用路径,而非完整类定义。你训练时BertModel是从第三方库导入的,但加载时Python在__main__模块中找不到该类的定义,因此抛出AttributeError。而列表、字典这类内置类型的类定义在固定系统模块中,所以pickle可以正常处理。
解决方案
1. 官方推荐:用state_dict保存/加载模型
PyTorch官方不建议直接pickle整个模型,保存模型的状态字典(仅保存权重参数)是更可靠的方式:
- 保存代码:
import torch torch.save(model.state_dict(), 'ner_bert_state_dict.pt')
- 加载代码:
from transformers import BertForTokenClassification import torch # 先初始化和训练时结构完全一致的模型实例 model = BertForTokenClassification.from_pretrained( 'bert-base-chinese', # 替换为你训练时用的预训练模型名 num_labels=你的标签数量 # 替换为实际标签数 ) # 加载权重 model.load_state_dict(torch.load('ner_bert_state_dict.pt')) model.eval() # 切换到评估模式
2. 临时兼容:修改模块引用(不推荐)
如果已经用pickle保存了模型,可在加载前将BertModel绑定到__main__模块,让pickle能找到类定义:
from transformers import BertModel import pickle import __main__ # 将BertModel注册到__main__模块 __main__.BertModel = BertModel # 加载模型 with open('model_pkl', 'rb') as file: model = pickle.load(file)
注意:这种方法依赖训练时的模块结构,换环境或修改类定义后极易失效,仅作临时救急使用。
内容的提问来源于stack exchange,提问作者majid bhatti
相关产品推荐
相关产品推荐

