直接修改模型最后一层需先确认其命名和结构,如ResNet用fc、VGG用classifier、ViT用heads;获取原in_features后新建nn.Linear(in_features, num_classes),并重置参数以避免残留旧值。

PyTorch预训练模型最后一层怎么替换成自定义分类数
直接改 model.classifier[-1] 或 model.fc 就行,但必须先确认原模型结构——不同 backbone 的输出层命名和位置差异很大,硬套会报 AttributeError: 'Sequential' object has no attribute 'fc' 这类错。
实操建议:
- 用
print(model)或list(model.children())[-1]查最后一级模块名,比如resnet50是fc,vgg16是classifier,vit_b_16是heads - 替换前记录原输出维度:
in_features = model.fc.in_features(ResNet)或model.classifier[-1].in_features(VGG),这是新层的in_features依据 - 新层必须用
nn.Linear(in_features, num_classes),不能漏掉num_classes—— 这是你要的分类数,不是原数据集的类别数
替换后权重初始化要不要重置
要。新层参数默认已随机初始化,但如果你手动创建了 nn.Linear 却没调用 reset_parameters(),可能残留旧值(尤其从 checkpoint 加载后部分替换时)。
常见疏漏点:
立即学习“Python免费学习笔记(深入)”;
- 只改了结构,没清缓存:训练前务必调用
model.apply(lambda m: m.reset_parameters() if isinstance(m, nn.Linear) else None),或者更稳妥地只重置你刚换的新层 - 用了
nn.Sequential包裹新层但忘了给它起名,导致无法单独 reset —— 建议显式赋值:model.fc = nn.Linear(2048, 12),而非model.classifier = nn.Sequential(..., new_layer) - ViT 类模型(如
vit_b_16)的heads是nn.Sequential(nn.Linear(...), nn.Dropout(), nn.Linear(...)),只换最后那个nn.Linear即可,别把整个heads替掉
动态映射分类数时如何避免 forward 报错
核心是确保新层输入特征维度与上游输出一致。典型错误是修改 num_classes 后忘记同步调整 in_features,导致 mat1 and mat2 shapes cannot be multiplied。
检查步骤:
- 打印上游模块输出 shape:
with torch.no_grad(): print(model.features(torch.randn(1,3,224,224)).shape)(对 CNN)或model.encoder(torch.randn(1,3,224,224)).shape(ViT),确认最后一维是否等于你设的in_features - 若用
torchvision.models.get_model('resnet50', weights='DEFAULT'),注意 0.14+ 版本返回的是带预处理的 wrapper,需访问.model属性才能拿到纯模型 - 迁移学习微调时,冻结参数(
requires_grad = False)只需加在特征提取部分,新分类层必须保持True,否则 loss 不下降
多任务或增量学习场景下怎么安全替换
不能直接覆盖原层变量名,否则旧层参数仍留在 model.state_dict() 里,保存加载时会冲突。
正确做法:
- 先删旧键:
state_dict = model.state_dict(); state_dict.pop('fc.weight', None); state_dict.pop('fc.bias', None),再 load 新权重 - 用
strict=False加载 checkpoint:model.load_state_dict(checkpoint, strict=False),缺失的 key(新层)自动跳过,多出的 key(旧层)也忽略 - 如果新增类别数比原来少(比如从1000→10),不要复用原
fc.weight[:10]—— 初始化更稳;若想保留语义相似性,可用 PCA 降维后再初始化,但非常规操作
fc,结果在 EfficientNet 上死磕 model.fc 却找不到。动手前花 10 秒 print(model) 看一眼,比查文档快。


















