使用torch.load加载PyTorch保存的BERT模型时出现transformers模块缺失报错
问题原因
这个报错的核心是你使用了torch.save(model, save_path)这种直接序列化整个模型对象的保存方式:
- transformers在v3.x和v4.x两个大版本之间做了模块路径重构:v3.x版本中BERT相关类的存放路径是
transformers.modeling_bert,v4.x版本重构后迁移到了transformers.models.bert.modeling_bert路径下 - 直接保存整个模型时,会把模型类的导入路径一起序列化,加载时必须和保存模型时用的transformers版本、模块路径完全一致才能正常反序列化。你两个环境的transformers版本和你保存模型时的版本都不匹配,所以分别触发了两种路径不存在的报错。
解决方案
分两种场景处理:
场景1:可以确定保存模型时的transformers版本
- 安装和保存时完全一致的transformers版本
- 加载模型后立刻导出权重文件,之后都用权重文件的方式加载,避免后续再出现同类问题:
# 安装对应版本后先加载完整模型 model = torch.load(r'Models-BERT\model-duplicates dropped(sampled)_v2', map_location=torch.device('cpu')) # 仅保存权重 torch.save(model.state_dict(), 'model_bert_weights.pt')
后续加载通用流程,不受版本变动影响:
from transformers import BertForSequenceClassification # 根据你实际的模型类调整 # 先实例化和训练时结构完全一致的模型 model = BertForSequenceClassification.from_pretrained('bert-base-chinese') # 初始化参数和训练时保持一致 # 加载权重 model.load_state_dict(torch.load('model_bert_weights.pt', map_location='cpu'))
场景2:无法确定保存模型时的版本
可以通过临时映射模块路径的方式绕过报错,根据你当前用的transformers版本选择对应代码:
当前用transformers 4.x版本,报错No module named 'transformers.modeling_bert'
加载前先添加模块路径映射:
import sys from transformers.models.bert import modeling_bert # 把旧版本的路径映射到当前版本的实际路径 sys.modules['transformers.modeling_bert'] = modeling_bert # 再执行加载代码 model = torch.load(r'Models-BERT\model-duplicates dropped(sampled)_v2', map_location=torch.device('cpu'))
当前用transformers 3.x版本,报错No module named 'transformers.models'
加载前先添加模块路径映射:
import sys from transformers import modeling_bert # 把新版本的路径映射到当前版本的实际路径 sys.modules['transformers.models.bert.modeling_bert'] = modeling_bert # 再执行加载代码 model = torch.load(r'Models-BERT\model-duplicates dropped(sampled)_v2', map_location=torch.device('cpu'))
注:如果使用的不是BERT而是其他类型的模型(比如RoBERTa、GPT2等),把上面代码中的
bert和modeling_bert替换为对应模型的名称即可。
内容的提问来源于stack exchange,提问作者Saksham Dubey
相关产品推荐
相关产品推荐

