PyTorch加载训练好的MobileNetV2模型出现ConvBNReLU属性错误如何解决?
错误原因
- 模型保存方式不规范:你训练时使用了
torch.save(model, 存储路径)直接序列化整个模型实例,该方式高度依赖训练环境的torchvision版本。不同版本torchvision中MobileNet系列的内部类ConvBNReLU存储路径发生了变动,加载环境和训练环境版本不一致时,反序列化就找不到对应类触发报错。 - 加载路径存在错误可能:你写的预期加载路径是MobileNetV2的模型路径,但报错日志显示实际运行时加载的是
models/MobileNetV3_CLEF_Small10/MobileNetV3_model_29.pt,先确认是否填错了加载路径。
修复方案
临时修复(可快速解决当前加载问题)
方案1(优先推荐):先初始化模型结构再加载权重
这是PyTorch官方推荐的加载方式,完全规避环境类依赖问题,代码示例如下:
import torch from torchvision import models from torch import nn device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") num_classes = 10 # 替换为你训练时实际的分类数量 # 初始化和训练时完全一致的模型结构 model = models.mobilenet_v2(pretrained=False) model.classifier = nn.Sequential( nn.Dropout(0.2), nn.Linear(model.last_channel, num_classes), ) model = model.to(device) # 加载已存储的模型数据,提取权重赋值给初始化好的结构 loaded_obj = torch.load("models/MobileNetV2_CLEF_Small10/MobileNetV2_model_29.pt", map_location=device) if isinstance(loaded_obj, torch.nn.Module): state_dict = loaded_obj.state_dict() else: state_dict = loaded_obj model.load_state_dict(state_dict)
方案2:补全类导入
如果不想修改现有加载逻辑,可以在加载代码的最开头添加对应类的导入,让反序列化时可以找到ConvBNReLU:
# 根据你当前环境的torchvision版本选一行执行 from torchvision.models.mobilenetv2 import ConvBNReLU # 适配0.11及以上版本torchvision # from torchvision.models.mobilenet import ConvBNReLU # 适配0.10及以下版本torchvision # 之后执行你原来的加载代码即可 model = torch.load("models/MobileNetV2_CLEF_Small10/MobileNetV2_model_29.pt")
长期规范
后续训练完模型保存时,仅保存权重参数,不要保存整个模型实例,从根源避免版本依赖问题:
# 训练完成后用该代码保存 torch.save(model.state_dict(), "你的模型存储路径.pt")
内容的提问来源于stack exchange,提问作者frex
相关产品推荐
相关产品推荐

