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

加载PyTorch模型时出现Missing key(s) in state_dict错误求助

模型加载时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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 16:24:57