PyTorch 2.0.0中InceptionV3 aux_logits参数报错及功能疑问
问题解答
1. 为啥会报"aux_logits expected value True but got False"?
在PyTorch 2.0.0及之后版本里,你用models.Inception_V3_Weights.DEFAULT加载的预训练权重,是基于**开启辅助分类器(aux_logits=True)**的InceptionV3结构训练出来的。要是直接把aux_logits设为False,模型结构和预训练权重的结构对不上,就会触发这个参数不兼容的报错。说白了就是:预训练权重对应的模型带辅助分类器,你不能硬套一个不带辅助分类器的结构去加载它。
2. aux_logits到底是干啥的?
没错,它就是和InceptionV3的辅助分类器绑定的:
- InceptionV3是深层网络,训练时容易出现梯度消失的问题。辅助分类器是从网络中间层(Inception模块之后)分支出来的小型分类器,训练阶段会和主分类器一起计算损失,帮梯度更顺畅地传到网络浅层,提升训练的稳定性和收敛速度。
- 到了推理阶段,辅助分类器没用,一般会关掉它(设
aux_logits=False),只留主分类器的输出。
3. 改代码解决报错的方案
你得先以aux_logits=True加载预训练模型,再手动关闭辅助分类器,同时替换全连接层:
def __init__(self, embed_size, trainCNN=False): super(encoderCNN, self).__init__() self.trainCNN = trainCNN # 先按aux_logits=True加载,匹配预训练权重的结构 self.inception = models.inception_v3(weights=models.Inception_V3_Weights.DEFAULT, aux_logits=True) # 关闭辅助分类器 self.inception.aux_logits = False # 移除辅助分类器的参数(可选,避免冗余参数占用资源) if hasattr(self.inception, 'AuxLogits'): self.inception.AuxLogits = None # 替换主分类器的全连接层 self.inception.fc = nn.Linear(self.inception.fc.in_features, embed_size) self.dropout= nn.Dropout(0.5) self.relu = nn.ReLU()
要是不需要训练CNN部分,还可以加载后冻结参数:
if not self.trainCNN: for param in self.inception.parameters(): param.requires_grad = False
内容的提问来源于stack exchange,提问作者MaheenUnzeelah
相关产品推荐
相关产品推荐

