如何在冻结特征提取层参数的前提下修改PyTorch预训练模型的类别数
嘿,我来帮你搞定这个问题!你的思路完全正确——冻结预训练模型的特征提取层,只训练新的分类头,这是迁移学习里高效又实用的做法。先给你梳理通用步骤,再针对你提到的每个模型给出具体代码实现,包括你已经写了的ResNet18,我也会帮你优化下之前的写法。
通用核心思路
- 加载预训练模型,用
pretrained=True(PyTorch 1.10+更推荐用weights=models.XXX_Weights.DEFAULT,不过为了兼容旧版本,下面还是用你习惯的写法) - 冻结特征层参数:遍历模型所有参数,设置
param.requires_grad = False,让这些层在训练时不更新权重 - 替换/修改分类头部:不同模型的分类头命名和结构差异很大,比如ResNet是
fc,VGG是classifier,需要针对性调整 - 新分类头的参数默认
requires_grad=True,不用额外设置,训练时会自动更新
1. ResNet18
你之前的写法方向对,但有个小问题:自定义的MyResModel没有包含ResNet的特征提取部分,直接传图像张量会报错。这里给你两种可行写法:
方式一:直接替换fc层(简洁版)
import torch import torch.nn as nn from torchvision import models # 加载预训练模型 resnet18 = models.resnet18(pretrained=True) # 冻结所有特征层参数 for param in resnet18.parameters(): param.requires_grad_(False) # 替换分类头:ResNet最后一层是fc,输入维度固定为512,输出设为你的类别数(比如3类) resnet18.fc = nn.Sequential( nn.Linear(512, 256), nn.ReLU(), nn.Dropout(p=0.5), nn.Linear(256, 3) )
方式二:自定义完整模型(修正你原来的写法)
import torch import torch.nn as nn from torchvision import models class MyResModel(torch.nn.Module): def __init__(self, pretrained_model, num_classes=3): super(MyResModel, self).__init__() # 保留预训练模型的特征提取部分(去掉最后一层fc) self.features = nn.Sequential(*list(pretrained_model.children())[:-1]) # 冻结特征层 for param in self.features.parameters(): param.requires_grad_(False) # 自定义分类头 self.classifier = nn.Sequential( nn.Linear(512, 256), nn.ReLU(), nn.Dropout(p=0.5), nn.Linear(256, num_classes), ) def forward(self, x): # 先过特征提取层得到512维特征 x = self.features(x) # 展平特征张量 x = torch.flatten(x, 1) # 再过分类头输出结果 x = self.classifier(x) return x # 使用方式 resnet18_pretrained = models.resnet18(pretrained=True) my_resnet = MyResModel(resnet18_pretrained, num_classes=3)
2. DenseNet161
DenseNet的分类头是classifier,输入维度可以通过densenet161.classifier.in_features获取(DenseNet161是2208):
densenet161 = models.densenet161(pretrained=True) # 冻结特征层 for param in densenet161.parameters(): param.requires_grad_(False) num_classes = 3 # 修改分类头 densenet161.classifier = nn.Sequential( nn.Linear(densenet161.classifier.in_features, 256), nn.ReLU(), nn.Dropout(0.5), nn.Linear(256, num_classes) )
3. Inception_v3
Inception_v3比较特殊,训练时会输出主结果和辅助结果,对应的分类头分别是fc和AuxLogits.fc。如果不需要辅助输出,只修改主分类头即可;如果要保留辅助训练,两个头都要改:
inception_v3 = models.inception_v3(pretrained=True) # 冻结特征层 for param in inception_v3.parameters(): param.requires_grad_(False) num_classes = 3 # 修改主分类头 inception_v3.fc = nn.Linear(inception_v3.fc.in_features, num_classes) # 可选:修改辅助分类头(训练时能帮助模型更快收敛) inception_v3.AuxLogits.fc = nn.Linear(inception_v3.AuxLogits.fc.in_features, num_classes)
注意:Inception_v3要求输入图像尺寸为(299,299),训练时记得调整输入大小;forward会返回两个输出,损失函数要同时处理主输出和辅助输出。
4. ShuffleNet_v2_x1_0
ShuffleNetV2的分类头是fc,输入维度为1024:
shufflenet_v2_x1_0 = models.shufflenet_v2_x1_0(pretrained=True) # 冻结特征层 for param in shufflenet_v2_x1_0.parameters(): param.requires_grad_(False) num_classes = 3 shufflenet_v2_x1_0.fc = nn.Sequential( nn.Linear(shufflenet_v2_x1_0.fc.in_features, 256), nn.ReLU(), nn.Dropout(0.5), nn.Linear(256, num_classes) )
5. MobileNet_v3_large / MobileNet_v3_small
MobileNetV3的分类头是classifier,是一个包含Dropout和Linear的Sequential层,直接替换整个classifier即可:
# MobileNet_v3_large mobilenet_v3_large = models.mobilenet_v3_large(pretrained=True) # 冻结特征层 for param in mobilenet_v3_large.parameters(): param.requires_grad_(False) num_classes = 3 mobilenet_v3_large.classifier = nn.Sequential( nn.Dropout(p=0.2), nn.Linear(mobilenet_v3_large.classifier[1].in_features, 256), nn.ReLU(), nn.Dropout(0.5), nn.Linear(256, num_classes) ) # MobileNet_v3_small同理 mobilenet_v3_small = models.mobilenet_v3_small(pretrained=True) for param in mobilenet_v3_small.parameters(): param.requires_grad_(False) mobilenet_v3_small.classifier = nn.Sequential( nn.Dropout(p=0.2), nn.Linear(mobilenet_v3_small.classifier[1].in_features, 128), nn.ReLU(), nn.Dropout(0.5), nn.Linear(128, num_classes) )
6. MNASNet1_0
MNASNet的分类头是classifier,输入维度为1280:
mnasnet1_0 = models.mnasnet1_0(pretrained=True) # 冻结特征层 for param in mnasnet1_0.parameters(): param.requires_grad_(False) num_classes = 3 mnasnet1_0.classifier = nn.Sequential( nn.Dropout(p=0.2), nn.Linear(mnasnet1_0.classifier[1].in_features, 256), nn.ReLU(), nn.Dropout(0.5), nn.Linear(256, num_classes) )
7. ResNeXt50_32x4d
ResNeXt和ResNet结构一致,分类头是fc,输入维度为512:
resnext50_32x4d = models.resnext50_32x4d(pretrained=True) # 冻结特征层 for param in resnext50_32x4d.parameters(): param.requires_grad_(False) num_classes = 3 resnext50_32x4d.fc = nn.Sequential( nn.Linear(512, 256), nn.ReLU(), nn.Dropout(0.5), nn.Linear(256, num_classes) )
8. VGG16
VGG的分类头是classifier,是一个包含多个Linear层的Sequential。有两种修改方式:
vgg16 = models.vgg16(pretrained=True) # 冻结特征层 for param in vgg16.parameters(): param.requires_grad_(False) num_classes = 3 # 方式一:替换整个分类头(完全自定义) vgg16.classifier = nn.Sequential( nn.Linear(25088, 4096), nn.ReLU(), nn.Dropout(0.5), nn.Linear(4096, 256), nn.ReLU(), nn.Dropout(0.5), nn.Linear(256, num_classes) ) # 方式二:只修改最后一层(保留预训练的分类头上层参数,更高效) vgg16.classifier[6] = nn.Linear(4096, num_classes)
额外小贴士
- 训练时只需要优化新分类头的参数,比如用
optimizer = torch.optim.Adam(model.fc.parameters(), lr=1e-3)(对应不同模型的分类头,比如VGG是model.classifier.parameters()) - 如果后续想微调特征层,可以解冻部分层(设置
param.requires_grad = True),然后用更小的学习率训练 - PyTorch 1.10+版本推荐用
weights参数替代pretrained,比如models.resnet18(weights=models.ResNet18_Weights.DEFAULT),能获取最新的预训练权重
内容的提问来源于stack exchange,提问作者Tim

