迁移学习后如何正确加载训练好的VGG模型并查看权重值
迁移学习VGG模型权重加载问题解答
问题背景
训练迁移学习VGG模型后,使用如下代码保存权重:
torch.save(vgg_based.state_dict(), 'model1.pth')
初次加载时直接实例化原生VGG19模型调用load_state_dict,出现RuntimeError:
RuntimeError: Error(s) in loading state_dict for VGG:
size mismatch for classifier.6.weight: copying a param with shape torch.Size([2, 4096]) from checkpoint, the shape in current model is torch.Size([1000, 4096]).
size mismatch for classifier.6.bias: copying a param with shape torch.Size([2]) from checkpoint, the shape in current model is torch.Size([1000])
报错原因是原生VGG19默认输出1000类,微调时把最后一层改成了2分类输出,结构不匹配导致加载失败。
加载代码正确性判断
你编写的复现分类头修改逻辑的加载代码是正确的,前提是满足两个条件:
- 训练阶段修改VGG19分类头的逻辑,和加载时写的修改逻辑完全一致(即都是去掉原最后一层1000分类全连接,替换为对应类别数的全连接层)
- 加载代码中
len(class_names)的数值和训练时的类别数完全一致(从报错看是2分类任务,只要len(class_names)=2就不会出现尺寸不匹配问题)
如果训练时冻结了VGG的特征提取层只训练分类头,这段代码加载后,特征层权重会匹配训练时用的基模型权重,分类层权重会匹配微调后的训练结果。
验证权重正确加载的方法
- 检查加载返回的键匹配结果
接收load_state_dict的返回值,打印缺失键和多余键,两个列表都为空就代表权重完全匹配:load_info = model.load_state_dict(torch.load('model1.pth')) print("未匹配到的缺失键:", load_info.missing_keys) print("权重文件中多余的未加载键:", load_info.unexpected_keys) - 直接对比权重数值
单独读取保存的权重字典,和模型当前层的权重形状、数值做比对:
形状一致、数值完全相同就代表权重加载正确。saved_state_dict = torch.load('model1.pth') # 对比最后分类层的形状 print("权重文件中classifier.6层权重形状:", saved_state_dict['classifier.6.weight'].shape) print("当前模型classifier.6层权重形状:", model.classifier[6].weight.shape) # 对比具体数值,取前5个值做快速校验 print("权重文件中该层权重前5个值:", saved_state_dict['classifier.6.weight'][0, :5]) print("当前模型该层权重前5个值:", model.classifier[6].weight[0, :5]) - 推理结果校验
用训练/验证集中标注正确的样本送入模型做推理,如果推理准确率和训练结束时的验证准确率基本一致,也可以确认模型加载无误。
内容的提问来源于stack exchange,提问作者user42
相关产品推荐
相关产品推荐

