You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

迁移学习后如何正确加载训练好的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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.27 11:18:16