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

如何在冻结特征提取层参数的前提下修改PyTorch预训练模型的类别数

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.30 09:02:41