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

加载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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 11:35:40