基于预训练DenseNet121修改分类器报错:矩阵维度不匹配
问题描述
在PyTorch中使用预训练DenseNet121搭建CNN时,尝试两种替换分类器的方式:第一种可正常运行,第二种为后续添加注意力机制而封装成自定义Module的方式,训练时报错:RuntimeError: mat1 and mat2 shapes cannot be multiplied (229376x7 and 1024x512)。两种方式的分类器结构一致,训练代码相同(batch size=32),为何第二种报错?
正常运行的第一种代码
import torch import torch.nn as nn from torchvision import models from collections import OrderedDict model = models.densenet121(pretrained=True) for param in model.parameters(): param.requires_grad = False classifier = nn.Sequential(OrderedDict([ ('fc1', nn.Linear(1024, 512)), ('relu', nn.ReLU()), ('fc2', nn.Linear(512, 10)), ('output', nn.LogSoftmax(dim=1)) ])) model.classifier = classifier
报错的第二种代码
import torch import torch.nn as nn from torchvision import models from collections import OrderedDict net = models.densenet121(pretrained=True) for param in net.parameters(): param.requires_grad = False class AttnDenseNet121(nn.Module): def __init__(self, num_classes, normalize_attn=False, dropout=None): super(AttnDenseNet121, self).__init__() self.features = net.features self.classifier = nn.Sequential(OrderedDict([ ('fc1', nn.Linear(1024, 512)), ('relu', nn.ReLU()), ('fc2', nn.Linear(512, 10)), ('output', nn.LogSoftmax(dim=1)) ])) def forward(self, x): x = self.features(x) out = self.classifier(x) return out model = AttnDenseNet121(num_classes=10, normalize_attn=True)
问题原因与解决方法
核心原因
原DenseNet的完整forward流程自带全局平均池化(AdaptiveAvgPool2d)和展平操作,而第二种自定义Module的forward函数中,仅调用了self.features(x),直接把特征图传入分类器,缺少了池化和展平步骤,导致输入分类器的张量形状不匹配:
- 第一种方式中,替换
model.classifier后,原模型的forward流程保持完整:features输出(batch_size,1024,7,7)的特征图 → 自动经过AdaptiveAvgPool2d((1,1))压缩为(batch_size,1024,1,1)→ 展平为(batch_size,1024)→ 传入分类器,维度匹配。 - 第二种方式中,
self.features(x)输出的(32,1024,7,7)张量直接进入Linear层,PyTorch会自动展平后两个维度,变成(32,229376),和Linear层要求的输入维度1024完全不匹配,因此触发形状错误。
解决方法
在自定义模型的forward函数中手动添加全局平均池化和展平操作,对齐原DenseNet的流程:
修正后的第二种代码:
import torch import torch.nn as nn from torchvision import models from collections import OrderedDict net = models.densenet121(pretrained=True) for param in net.parameters(): param.requires_grad = False class AttnDenseNet121(nn.Module): def __init__(self, num_classes, normalize_attn=False, dropout=None): super(AttnDenseNet121, self).__init__() self.features = net.features # 新增全局平均池化层,与原DenseNet一致 self.avgpool = nn.AdaptiveAvgPool2d((1, 1)) self.classifier = nn.Sequential(OrderedDict([ ('fc1', nn.Linear(1024, 512)), ('relu', nn.ReLU()), ('fc2', nn.Linear(512, 10)), ('output', nn.LogSoftmax(dim=1)) ])) def forward(self, x): x = self.features(x) # 添加池化和展平步骤 x = self.avgpool(x) x = torch.flatten(x, 1) out = self.classifier(x) return out model = AttnDenseNet121(num_classes=10, normalize_attn=True)
内容的提问来源于stack exchange,提问作者Ayush
相关产品推荐
相关产品推荐

