PyTorch预训练ResNet如何启用Bias项并解决state_dict加载报错
报错原因
PyTorch官方实现的ResNet所有卷积层默认设置bias=False,因为卷积层后紧跟BatchNorm层时,卷积的bias会在BN归一化计算时被完全抵消,对输出没有实际贡献,所以官方发布的预训练权重里根本没有存储卷积层对应的bias参数。你修改自定义ResNet结构把卷积层bias设为True后,直接加载预训练权重就会触发缺失key的报错。
解决步骤
- 第一步:实例化你修改好的、所有卷积层开启bias的自定义ResNet结构,加载时不要开预训练选项,同时把最后的全连接层替换成你任务对应类别数的线性层。
示例代码:# 此处resnet18是你参考官方源码修改了卷积层bias=True的自定义实现 net = resnet18(pretrained=False) net.fc = nn.Linear(512, num_classes) - 第二步:加载官方原始结构的预训练ResNet,取出它的权重字典,和你自定义网络的权重字典做匹配过滤,只保留键名一致、参数形状一致的权重,过滤掉不存在的bias项、形状不匹配的fc层参数。
示例代码:import torch import torchvision.models as models # 加载官方原始结构的预训练ResNet18,提取预训练权重 official_pretrained = models.resnet18(pretrained=True) pretrained_params = official_pretrained.state_dict() # 提取自定义网络的当前参数 custom_params = net.state_dict() # 筛选出两边匹配的参数 matched_params = {} for k, v in pretrained_params.items(): if k in custom_params and v.shape == custom_params[k].shape: matched_params[k] = v # 用匹配到的预训练参数覆盖自定义网络的对应参数 custom_params.update(matched_params) net.load_state_dict(custom_params) - 第三步:直接正常使用网络即可。加载过程不会再报错,所有和预训练结构匹配的卷积核、BN层参数都会正常加载,新增的卷积层bias参数会使用PyTorch默认的层初始化规则随机初始化,在后续训练中更新。
注意:如果你的自定义网络里卷积层后面仍然保留了BatchNorm层,开启卷积bias不会带来实际的效果提升,只会额外增加参数量,做预训练/非预训练对比实验时要注意控制变量,避免无关变量影响实验结论。
内容的提问来源于stack exchange,提问作者Siddeshwar Raghavan
相关产品推荐
相关产品推荐

