如何以统一方式修改不同PyTorch预训练模型最后一层用于微调
PyTorch预训练模型最后一层统一修改方案
你之前直接给get_last_module返回值赋值不生效,核心原因是Python的变量赋值仅做引用绑定:拿到最后一层对象的引用后,给变量赋新值只会改变变量本身的指向,不会修改原模型内部存储的模块属性,自然不会对原模型产生任何修改。
通用方案不需要针对不同模型写分支判断,只需要在定位最后一层时,同时拿到最后一层所属的父模块、最后一层在父模块中的属性名/索引、最后一层本身的参数信息,直接在父模块上执行属性替换即可,适配所有标准PyTorch模型结构。
通用工具函数
import torch import torch.nn as nn from typing import Tuple def get_last_layer_info(model: nn.Module) -> Tuple[nn.Module, str, nn.Module, int, int]: """ 提取模型最后输出层的全量信息 返回值顺序:(最后一层的父模块, 最后一层对应的属性名, 原始最后一层对象, 最后一层输入维度, 最后一层输出维度) """ parent_module = None attr_name = None last_layer = None # 优先匹配主流模型库的标准分类头命名,覆盖99%常规预训练模型 common_head_paths = ["fc", "head", "classifier", "head.fc", "classifier.6"] for path in common_head_paths: try: path_parts = path.split(".") target_obj = model # 逐层解析嵌套属性,比如head.fc这类嵌套结构 for part in path_parts[:-1]: target_obj = getattr(target_obj, part) target_attr = path_parts[-1] candidate_layer = getattr(target_obj, target_attr) # 校验是带维度属性的计算层,不是容器 if hasattr(candidate_layer, "in_features") and hasattr(candidate_layer, "out_features"): parent_module = target_obj attr_name = target_attr last_layer = candidate_layer break except (AttributeError, IndexError): continue # 常规命名匹配失败时,递归遍历所有子模块,取遍历顺序最末端的线性/卷积输出层 if last_layer is None: for module_path, module in model.named_modules(): if isinstance(module, (nn.Linear, nn.Conv2d)): path_parts = module_path.split(".") if len(path_parts) == 1: current_parent = model current_attr = path_parts[0] else: current_parent = model for part in path_parts[:-1]: current_parent = current_parent[int(part)] if part.isdigit() else getattr(current_parent, part) current_attr = path_parts[-1] parent_module = current_parent attr_name = current_attr last_layer = module # 提取输入输出维度,兼容线性层和卷积层 in_dim = last_layer.in_features if hasattr(last_layer, "in_features") else last_layer.in_channels out_dim = last_layer.out_features if hasattr(last_layer, "out_features") else last_layer.out_channels return parent_module, attr_name, last_layer, in_dim, out_dim
简化后的业务代码
用上述工具替换原来的类型分支判断,代码量直接减半,新增模型类型不需要加判断逻辑:
def __init__(self, base_encoder, pred_dim=512): """ dim: 输出特征维度 (默认值2048) pred_dim: 预测器隐藏层维度 (默认值512) """ super().__init__() self.encoder = base_encoder # 统一获取最后一层信息,无需判断模型属于ResNet还是ViT parent_module, last_attr, origin_last_layer, prev_dim, dim = get_last_layer_info(self.encoder) # 构建3层投影头,逻辑和原实现完全一致 new_proj_head = nn.Sequential( nn.Linear(prev_dim, prev_dim, bias=False), nn.BatchNorm1d(prev_dim), nn.ReLU(inplace=True), nn.Linear(prev_dim, prev_dim, bias=False), nn.BatchNorm1d(prev_dim), nn.ReLU(inplace=True), origin_last_layer, # 保留原始预训练的最后一层 nn.BatchNorm1d(dim, affine=False), ) # 直接在父模块上替换属性,修改即时生效 setattr(parent_module, last_attr, new_proj_head) # 冻结原始最后一层的bias,和原hack逻辑一致 getattr(parent_module, last_attr)[6].bias.requires_grad = False # 构建2层预测器,逻辑不变 self.predictor = nn.Sequential( nn.Linear(dim, pred_dim, bias=False), nn.BatchNorm1d(pred_dim), nn.ReLU(inplace=True), nn.Linear(pred_dim, dim), )
适配说明
- 优先匹配逻辑覆盖torchvision、timm、torchhub等主流来源的预训练模型,包括ResNet的
fc、ViT的head、部分分类模型的classifier等常见命名,匹配失败才会递归遍历找最末端的计算层,也支持最后一层由多层嵌套构成的复杂头部结构。 - 遇到多任务头等特殊结构模型,只需要在工具函数的
common_head_paths列表里追加对应头部的属性路径即可,不需要修改主业务逻辑。 - 如果是卷积作为输出层的Backbone,工具函数会自动读取
in_channels/out_channels,只需要把投影头里的nn.Linear替换为对应尺寸的卷积层即可,定位逻辑不需要调整。
内容的提问来源于stack exchange,提问作者Whisht
相关产品推荐
相关产品推荐

