使用efficient_b4做迁移学习时如何正确获取全连接层特征数num_ftrs
问题原因
PyTorch官方实现的EfficientNet系列模型的分类层命名和ResNet等常用模型不同,没有名为fc的属性,它的分类头命名为classifier。
解决方法
直接读取classifier序列中全连接层的输入特征数即可,正确代码如下:
import torchvision.models as models model = models.efficientnet_b4(pretrained = True) # EfficientNet的classifier由[Dropout, Linear]两层组成,索引为1的是全连接层 num_ftrs = model.classifier[1].in_features
如果后续要替换全连接层适配自己的分类任务,可参考如下写法(示例为10分类场景):
import torch model.classifier[1] = torch.nn.Linear(num_ftrs, 10)
补充说明
如果遇到其他不确定层命名的模型,可直接打印模型结构查看最后几层的命名和层级关系:
print(model)
内容的提问来源于stack exchange,提问作者Samriddha Majumdar
相关产品推荐
相关产品推荐

