PyTorch图像分类器模型加载报错:分类器state_dict不匹配求助
模型加载错误的原因与修复方案
看起来你遇到的问题核心是模型架构不匹配和分类器加载顺序错误导致的state_dict键不匹配,我来一步步帮你梳理解决:
问题1:模型架构不一致
你保存的checkpoint里明确标注了arch: 'vgg19',但加载时却初始化了models.vgg16(pretrained=True)——VGG16和VGG19的基础特征提取层结构不同,对应的state_dict键自然无法匹配,这是第一个致命问题。
问题2:分类器加载顺序颠倒
你自定义的分类器用了fc1、fc2这类命名的层,但VGG默认的分类器层是classifier.0、classifier.3、classifier.6。你现在的代码是先加载整个模型的state_dict,再替换分类器,这就导致加载时模型还是默认的分类器结构,和你保存的自定义分类器参数完全不兼容,从而出现"Missing key"和"Unexpected key"的错误。
修复后的加载函数
我调整了加载逻辑,解决了上述两个问题,代码如下:
def load_checkpoint(filepath): checkpoint = torch.load(filepath) # 动态加载与保存时一致的模型架构,避免硬编码错误 model = getattr(models, checkpoint['arch'])(pretrained=True) # 先替换为自定义分类器,再加载模型参数——这步顺序很关键! model.classifier = checkpoint['classifier'] model.load_state_dict(checkpoint['state_dict']) model.class_to_idx = checkpoint['class_to_idx'] model.epochs = checkpoint['epochs'] # 重新初始化optimizer,再加载其状态(必须先绑定当前模型参数) learn_rate = checkpoint['learn_rate'] momentum = checkpoint['momentum'] optimizer = torch.optim.SGD(model.classifier.parameters(), lr=learn_rate, momentum=momentum) optimizer.load_state_dict(checkpoint['optimizer']) return learn_rate, optimizer, model # 调用加载函数 learn_rate, optimizer, model = load_checkpoint('checkpoint.pth')
关键修复点说明
- 动态匹配模型架构:用
getattr(models, checkpoint['arch'])代替硬编码的models.vgg16,确保加载的模型和保存时完全一致,不管你后续换VGG11还是ResNet,都不会出架构不匹配的问题。 - 调整分类器加载顺序:先把模型的分类器替换为你保存的自定义版本,再加载整个模型的state_dict,这样state_dict里的分类器参数(
fc1、fc2)就能和当前模型的结构对应上,不会出现键不匹配的错误。 - 正确初始化优化器:优化器是和模型参数绑定的,加载时必须先基于当前模型的参数创建优化器,再加载它的state_dict,否则会出现参数不匹配的隐性问题。
额外建议
- 后续保存模型时,建议不要直接序列化
classifier对象,而是保存分类器的结构参数(比如输入维度、隐藏层大小、输出维度等),然后在加载时重新构建分类器,这样可以避免序列化对象带来的兼容性问题(比如不同PyTorch版本可能无法正常加载序列化的自定义模块)。 - 如果训练和加载的设备不同(比如训练用GPU,加载用CPU),记得在
torch.load时加上map_location参数:torch.load(filepath, map_location='cpu'),避免设备不兼容的错误。
内容的提问来源于stack exchange,提问作者rachelvsamuel
相关产品推荐
相关产品推荐

