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

PyTorch修改预训练Inception v3为多输出模型时维度报错如何解决

问题原因
  • 核心是InceptionV3默认开启辅助分类器(aux_logits=True),且直接用.children()切层的方式不符合InceptionV3的结构特性:
    1. 训练模式下开启aux_logits的InceptionV3前向传播返回的是(主输出, 辅助输出)的二元组,不是单个特征张量,直接传入后续池化层会导致维度异常
    2. 你当前的切层方式没有移除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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 12:36:04