加载PyTorch模型时出现Missing key(s) in state_dict错误求助
问题描述
我尝试加载模型,代码如下:
model = AlexNet3DDropoutRegression(9600) model_save_location = 'my_model.pt' model.load_state_dict(torch.load(model_save_location, map_location='cpu'))
此前保存模型的代码为:
torch.save(self.model.state_dict(), self.cli_args.model_save_location)
但加载时出现错误:
Missing key(s) in state_dict: "features.0.weight", "features.0.bias", "features.1.weight", "features.1.bias", "features.1.running_mean", "features.1.running_var", "features.4.weight", "features.4.bias", "features.5.weight", "features.5.bias", "features.5.running_mean", "features.5.running_var", "features.8.weight", "features.8.bias", "features.9.weight", "features.9.bias", "features.9.running_mean", "features.9.running_var", "features.11.weight", "features.11.bias", "features.12.weight", "features.12.bias", "features.12.running_mean", "features.12.running_var", "features.14.weight", "features.14.bias", "features.15.weight", "features.15.bias", "features.15.running_mean", "features.15.running_var", "classifier.1.weight", "classifier.1.bias", "classifier.4.weight", "classifier.4.bias".
Unexpected key(s) in state_dict: "module.features.0.weight", "module.features.0.bias", "module.features.1.weight", "module.features.1.bias", "module.features.1.running_mean", "module.features.1.running_var", "module.features.1.num_batches_tracked", "module.features.4.weight", "module.features.4.bias", "module.features.5.weight", "module.features.5.bias", "module.features.5.running_mean", "module.features.5.running_var", "module.features.5.num_batches_tracked", "module.features.8.weight", "module.features.8.bias", "module.features.9.weight", "module.features.9.bias", "module.features.9.running_mean", "module.features.9.running_var", "module.features.9.num_batches_tracked", "module.features.11.weight", "module.features.11.bias", "module.features.12.weight", "module.features.12.bias", "module.features.12.running_mean", "module.features.12.running_var", "module.features.12.num_batches_tracked", "module.features.14.weight", "module.features.14.bias", "module.features.15.weight", "module.features.15.bias", "module.features.15.running_mean", "module.features.15.running_var", "module.features.15.num_batches_tracked", "module.classifier.1.weight", "module.classifier.1.bias", "module.classifier.4.weight", "module.classifier.4.bias".
已确认保存与加载使用同一Python虚拟环境的PyTorch版本,完整的AlexNet3D模型代码如下:
import math import torch.nn as nn class AlexNet3D(nn.Module): def get_head(self): return nn.Sequential(nn.Dropout(), nn.Linear(self.input_size, 64), nn.ReLU(inplace=True), nn.Dropout(), nn.Linear(64, 1), ) def __init__(self, input_size): super().__init__() self.input_size = input_size self.features = nn.Sequential( nn.Conv3d(1, 64, kernel_size=(5, 5, 5), stride=(2, 2, 2), padding=0), nn.BatchNorm3d(64), nn.ReLU(inplace=True), nn.MaxPool3d(kernel_size=3, stride=3), nn.Conv3d(64, 128, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=0), nn.BatchNorm3d(128), nn.ReLU(inplace=True), nn.MaxPool3d(kernel_size=3, stride=3), nn.Conv3d(128, 192, kernel_size=(3, 3, 3), padding=1), nn.BatchNorm3d(192), nn.ReLU(inplace=True), nn.Conv3d(192, 192, kernel_size=(3, 3, 3), padding=1), nn.BatchNorm3d(192), nn.ReLU(inplace=True), nn.Conv3d(192, 128, kernel_size=(3, 3, 3), padding=1), nn.BatchNorm3d(128), nn.ReLU(inplace=True), nn.MaxPool3d(kernel_size=3, stride=3), ) self.classifier = self.get_head() for m in self.modules(): if isinstance(m, nn.Conv2d): n = m.kernel_size[0] * m.kernel_size[1] * m.out_channels m.weight.data.normal_(0, math.sqrt(2. / n)) elif isinstance(m, nn.BatchNorm3d): m.weight.data.fill_(1) m.bias.data.zero_() def forward(self, x): xp = self.features(x) x = xp.view(xp.size(0), -1) x = self.classifier(x) return [x, xp]
问题原因
错误提示里的module.xxx前缀说明:保存模型时,模型被nn.DataParallel或nn.parallel.DistributedDataParallel这类并行包装器包裹过,包装器会给所有模型参数的键加上module.前缀。而加载时直接初始化了原始的AlexNet3DDropoutRegression模型,它的参数键没有这个前缀,导致两者不匹配。
另外注意加载时用的是AlexNet3DDropoutRegression,但提供的模型代码是AlexNet3D,需要确保这两个模型的结构完全一致,不过当前核心问题是参数键的前缀差异。
解决方案
方案一:加载时去除参数键的module.前缀
修改加载代码,手动处理state_dict的键:
model = AlexNet3DDropoutRegression(9600) model_save_location = 'my_model.pt' state_dict = torch.load(model_save_location, map_location='cpu') # 去除所有键的module.前缀 new_state_dict = {k.replace('module.', ''): v for k, v in state_dict.items()} model.load_state_dict(new_state_dict)
方案二:加载时用并行包装器包裹模型
如果需要保持模型的并行结构,可在加载时用nn.DataParallel包裹模型:
model = AlexNet3DDropoutRegression(9600) model = nn.DataParallel(model) # 用DataParallel包裹模型 model_save_location = 'my_model.pt' model.load_state_dict(torch.load(model_save_location, map_location='cpu'))
额外注意
务必检查AlexNet3DDropoutRegression和保存时的AlexNet3D结构是否完全一致,包括层的数量、类型、输入输出维度,避免因模型结构差异导致其他键不匹配问题。
内容的提问来源于stack exchange,提问作者Paul Reiners

