如何在PyTorch中为修改后的预训练模型加载权重时使用strict=False
解决方案
你可以直接替换掉原来手动遍历权重的冗余逻辑,改用PyTorch原生的load_state_dict方法,这个方法本身就支持strict=False参数,完美适配你修改过模型结构的场景。修改后的完整代码如下:
import os import torch import torch.nn as nn class MyClass(nn.Module): def __init__(self, pretrained=False): super(MyClass, self).__init__() self.encoder = S3D_featureExtractor_multi_output() if pretrained: weight_dict = torch.load(os.path.join('models', 'weights.pt')) # 使用load_state_dict并设置strict=False missing_keys, unexpected_keys = self.encoder.load_state_dict(weight_dict, strict=False) # 可选:打印不匹配的键,方便排查 if missing_keys: print(f"模型中存在但权重里没有的层: {missing_keys}") if unexpected_keys: print(f"权重里存在但模型中没有的层: {unexpected_keys}") print('Loading finished!') def forward(self, x): a, b = self.encoder(x) return a, b
关键修改说明:
- 删掉了原来手动遍历权重字典、逐个复制张量的代码,这部分逻辑完全可以被
load_state_dict替代,更简洁高效 - 调用
self.encoder.load_state_dict(weight_dict, strict=False)时,strict=False会自动忽略两种不匹配情况:- 当前模型有但预训练权重里没有的层(missing_keys)
- 预训练权重里有但当前模型已经删掉/修改的层(unexpected_keys)
- 可选保留打印不匹配键的逻辑,方便你确认哪些层没有被加载,排查模型修改是否符合预期
内容的提问来源于stack exchange,提问作者dtr43
相关产品推荐
相关产品推荐

