微调自定义ResNet18架构后如何继续使用预训练权重无需从头训练
报错原因
你新增的layer_attend1层的参数不存在于旧检查点文件中,PyTorch默认的load_state_dict方法要求模型结构和检查点的参数键完全匹配才会执行加载,因此触发缺失键的报错。
核心问题解答
只要你没有修改ResNet18原有基础层的结构(卷积通道数、卷积核尺寸、层排列顺序等核心属性未改动),仅新增自定义层的情况下,完全可以正常设置pretrained=True使用预训练模式,官方预训练权重会正常加载到结构匹配的原有层上,新增层会走默认初始化逻辑,不会互相影响。
小幅改动后加载旧检查点的解决方案
方法1:开启非严格加载(最便捷)
给load_state_dict添加strict=False参数即可,PyTorch会自动忽略缺失的键和多余的键,原有结构匹配的参数会正常加载:
tnet.load_state_dict(checkpoint['state_dict'], strict=False)
说明:新增的自定义层会保持初始化状态,后续训练过程中会和其他参数一起更新,不需要从头训练整个模型。
方法2:手动过滤参数(更可控)
如果需要明确控制加载的参数,你可以手动过滤出检查点和当前模型匹配的参数后再更新:
# 获取当前模型的参数字典 model_dict = tnet.state_dict() # 筛选出检查点中存在且和当前模型匹配的参数 matched_pretrained_dict = {k: v for k, v in checkpoint['state_dict'].items() if k in model_dict} # 更新当前模型的参数字典 model_dict.update(matched_pretrained_dict) # 加载更新后的参数字典 tnet.load_state_dict(model_dict)
附加警告处理
你遇到的transforms.Scale弃用警告,只需将代码中所有的transforms.Scale替换为transforms.Resize即可消除。
内容的提问来源于stack exchange,提问作者Mona Jalal
相关产品推荐
相关产品推荐

