加载AudioClassifier模型state_dict时触发RuntimeError的问题咨询
模型state_dict加载失败的原因与解决方法
核心原因
训练时保存的模型结构和当前加载时使用的模型结构不一致:
- 训练时的模型没有第五个卷积块(
conv5、relu5、bn5),最后一层卷积输出通道数为64,因此全连接层lin的输入维度是64; - 当前加载时的模型新增了第五个卷积块,输出通道数变为128,全连接层输入维度也改成了128,导致state_dict中没有对应新增层的参数,同时全连接层参数形状不匹配。
解决方案
方案1:还原训练时的模型结构(推荐)
修改当前的AudioClassifier类,移除第五个卷积块,并将全连接层的输入改回64,和训练时的结构完全一致:
class AudioClassifier(nn.Module): def __init__(self): super().__init__() conv_layers = [] # 保留前四个卷积块 self.conv1 = nn.Conv2d(2, 8, kernel_size=(5, 5), stride=(2, 2), padding=(2, 2)) self.relu1 = nn.ReLU() self.bn1 = nn.BatchNorm2d(8) init.kaiming_normal_(self.conv1.weight, a=0.1) self.conv1.bias.data.zero_() conv_layers += [self.conv1, self.relu1, self.bn1] self.conv2 = nn.Conv2d(8, 16, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1)) self.relu2 = nn.ReLU() self.bn2 = nn.BatchNorm2d(16) init.kaiming_normal_(self.conv2.weight, a=0.1) self.conv2.bias.data.zero_() conv_layers += [self.conv2, self.relu2, self.bn2] self.conv3 = nn.Conv2d(16, 32, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1)) self.relu3 = nn.ReLU() self.bn3 = nn.BatchNorm2d(32) init.kaiming_normal_(self.conv3.weight, a=0.1) self.conv3.bias.data.zero_() conv_layers += [self.conv3, self.relu3, self.bn3] self.conv4 = nn.Conv2d(32, 64, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1)) self.relu4 = nn.ReLU() self.bn4 = nn.BatchNorm2d(64) init.kaiming_normal_(self.conv4.weight, a=0.1) self.conv4.bias.data.zero_() conv_layers += [self.conv4, self.relu4, self.bn4] # 移除第五个卷积块相关代码 self.ap = nn.AdaptiveAvgPool2d(output_size=1) # 全连接层输入改回64 self.lin = nn.Linear(in_features=64, out_features=10) self.conv = nn.Sequential(*conv_layers) def forward(self, x): x = self.conv(x) x = self.ap(x) x = x.view(x.shape[0], -1) x = self.lin(x) return x
之后执行加载代码即可正常运行:
model = AudioClassifier() model.load_state_dict(torch.load('model_v1.pth')) model.eval()
方案2:强制加载(不推荐,会损失性能)
如果必须使用当前带第五个卷积块的模型,可以设置strict=False跳过参数匹配检查,但新增的conv5和bn5会使用随机初始化的参数,模型性能会大幅下降:
model = AudioClassifier() # strict=False允许加载不匹配的参数 model.load_state_dict(torch.load('model_v1.pth'), strict=False) model.eval()
内容的提问来源于stack exchange,提问作者tjampolay
相关产品推荐
相关产品推荐

