PyTorch修改预训练Inception v3为多输出模型时维度报错如何解决
问题原因
- 核心是
InceptionV3默认开启辅助分类器(aux_logits=True),且直接用.children()切层的方式不符合InceptionV3的结构特性:- 训练模式下开启
aux_logits的InceptionV3前向传播返回的是(主输出, 辅助输出)的二元组,不是单个特征张量,直接传入后续池化层会导致维度异常 - 你当前的切层方式没有移除AuxLogits模块,输出的是2维的分类结果张量(维度[50, 1000]),不是4维的特征图,和报错信息完全匹配
- 训练模式下开启
- 补充:你的
__init__代码存在逻辑漏洞,仅当pretrained=True时才会定义self.model,如果传入pretrained=False会直接报属性不存在的错误
解决方法
修改模型定义即可,推荐两种修改方案:
方案1(更简单,直接替换原模型的全连接层)
import torch.nn as nn import torch.nn.functional as F from torchvision import models class CNN1(nn.Module): def __init__(self, pretrained): super(CNN1, self).__init__() # 实例化时直接关闭辅助分类器 self.model = models.inception_v3(pretrained=pretrained, aux_logits=False) # 获取原fc层的输入维度 in_features = self.model.fc.in_features # 删掉原有的fc层 self.model.fc = nn.Identity() self.fc0 = nn.Linear(in_features, 10) #digit 0 self.fc1 = nn.Linear(in_features, 10) #digit 1 self.fc2 = nn.Linear(in_features, 10) #digit 2 self.fc3 = nn.Linear(in_features, 10) #digit 3 def forward(self, x): # 关闭aux_logits后模型输出直接是铺平的特征张量,维度为[batch_size, 2048] x = self.model(x) label0 = self.fc0(x) label1 = self.fc1(x) label2= self.fc2(x) label3= self.fc3(x) return {'label0': label0, 'label1': label1,'label2':label2, 'label3': label3}
方案2(保留你原来的切层逻辑,修正错误点)
class CNN1(nn.Module): def __init__(self, pretrained): super(CNN1, self).__init__() # 不管pretrained取值都实例化模型,同时关闭aux_logits self.model = models.inception_v3(pretrained=pretrained, aux_logits=False) # 置空AuxLogits模块避免干扰 self.model.AuxLogits = None modules = list(self.model.children())[:-1] # delete the last fc layer. self.features = nn.Sequential(*modules) self.fc0 = nn.Linear(2048, 10) #digit 0 self.fc1 = nn.Linear(2048, 10) #digit 1 self.fc2 = nn.Linear(2048, 10) #digit 2 self.fc3 = nn.Linear(2048, 10) #digit 3 def forward(self, x): bs, _, _, _ = x.shape x = self.features(x) x = F.adaptive_avg_pool2d(x, 1).reshape(bs, -1) label0 = self.fc0(x) label1 = self.fc1(x) label2= self.fc2(x) label3= self.fc3(x) return {'label0': label0, 'label1': label1,'label2':label2, 'label3': label3}
额外注意
如果验证阶段也报类似维度错误,需要在验证前调用model.eval()切换到评估模式,评估模式下InceptionV3不会返回辅助输出。
内容的提问来源于stack exchange,提问作者piper11
相关产品推荐
相关产品推荐

