PyTorch保存的模型无论输入如何均返回相同预测的问题排查
问题诊断与解决方案
核心问题
基于VGG16训练笔记本4个部件(back/front/keyboard/monitor)的损伤分级模型,训练后即时测试结果正常,但保存再加载后,所有部件测试均返回S等级,问题源于模型保存与加载的逻辑错误。
错误点分析
1. 模型保存逻辑缺陷
- 仅保存单个模型:训练代码为每个部件训练独立模型并存储在
models_dict中,但保存代码仅保存了最后一个训练的模型(monitor部件),其余三个部件的模型未被保存。 - 无效代码破坏模型结构:保存代码中
model.fc = nn.Linear(num_features, 3)完全错误,VGG16无fc层,分类层为classifier,这行代码会篡改当前模型结构。 - 冗余参数存储:checkpoint中重复存储
classifier.6.weight/bias属于多余操作,直接保存model_state_dict即可完整记录模型权重。
2. 模型加载逻辑混乱
- 单模型复用错误:加载时用同一个模型对应所有部件,忽略了每个部件对应独立训练模型的事实。
- 权重加载逻辑冲突:手动赋值
classifier.6权重后又调用load_state_dict,且设置strict=False忽略键匹配错误,导致部分权重未正确加载。 - 模型初始化不一致:加载时用
pretrained=False初始化模型,特征提取层为随机权重,与训练时用预训练VGG16的逻辑不符,导致输出异常。
修正方案
方案1:批量保存所有模型
修改保存代码,将所有部件的模型一次性保存:
# 保存所有部件的模型字典 checkpoint_path = '/content/drive/MyDrive/vggnet_all_parts.pth' torch.save(models_dict, checkpoint_path)
对应加载代码:
import torch import torchvision.models as models device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # 加载所有部件的模型 models_dict = torch.load('/content/drive/MyDrive/vggnet_all_parts.pth') # 将模型移至设备并设置为评估模式 for part, model in models_dict.items(): model.to(device) model.eval()
方案2:单独保存每个模型
若需分部件管理模型,可单独保存每个模型:
# 逐个保存各部件模型 for part, model in models_dict.items(): checkpoint_path = f'/content/drive/MyDrive/vggnet_{part}.pth' torch.save(model.state_dict(), checkpoint_path)
对应加载代码:
import torch import torchvision.models as models from torch import nn device = torch.device("cuda" if torch.cuda.is_available() else "cpu") target_parts = ["back", "front", "keyboard", "monitor"] models_dict = {} for part in target_parts: # 构建与训练时一致的模型结构 model = models.vgg16(pretrained=True) num_features = model.classifier[6].in_features model.classifier[6] = nn.Linear(num_features, 3) # 加载对应部件的权重 checkpoint_path = f'/content/drive/MyDrive/vggnet_{part}.pth' model.load_state_dict(torch.load(checkpoint_path)) # 设置为评估模式并移至设备 model.to(device) model.eval() models_dict[part] = model
测试代码保持不变
原测试代码逻辑正确,只需确保models_dict中存储的是对应部件的正确模型即可。
额外注意事项
- 训练与加载的模型结构必须完全一致,包括预训练权重的使用。
- 保存模型优先选择保存
state_dict或完整模型,避免冗余参数。 - 加载后必须调用
model.eval()关闭 dropout 等训练模式特有的层。
内容的提问来源于stack exchange,提问作者강병국
相关产品推荐
相关产品推荐

