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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 22:06:26