PyTorch微调预训练MobileNet_V3_Large的线性层设置方法问询
MobileNet_V3_Large 微调分类层修改方法
PyTorch官方实现的MobileNet_V3_Large的分类头结构和MobileNet_V2有差异,替换最后一层线性层的方法如下:
首先给出完整示例代码:
import torch.nn as nn from torchvision import models # 加载预训练权重 # PyTorch 1.12及更早版本使用此行 model_ft = models.mobilenet_v3_large(pretrained=True, progress=True) # PyTorch 1.13及以上版本建议使用此行替换上方代码 # model_ft = models.mobilenet_v3_large(weights=models.MobileNet_V3_Large_Weights.DEFAULT, progress=True) # 获取分类头最后一层线性层的输入特征数 in_features = model_ft.classifier[3].in_features # 替换为适配自定义类别数的线性层 model_ft.classifier[3] = nn.Linear(in_features, out_features=len(class_names))
如果需要只训练分类头、冻结主干特征提取层,可以添加以下代码:
for param in model_ft.features.parameters(): param.requires_grad = False
通用排查方法
如果遇到其他MobileNet变体不知道怎么改分类层,可以直接打印分类头结构,确认需要替换的层的索引和输入维度:
print(model_ft.classifier)
以默认预训练MobileNet_V3_Large为例,打印输出如下,可见最后一层线性层索引为3:
Sequential( (0): Linear(in_features=960, out_features=1280, bias=True) (1): Hardswish() (2): Dropout(p=0.2, inplace=True) (3): Linear(in_features=1280, out_features=1000, bias=True) )
内容的提问来源于stack exchange,提问作者logankilpatrick
相关产品推荐
相关产品推荐

