运行XLSR PyTorch模型时state_dict键不匹配报错,求解决方案
问题描述
尝试运行某仓库中的PyTorch模型测试脚本,指定正确的best.pt路径后执行python test.py,出现模型加载报错。
测试脚本代码片段
device = 'cuda:' + str(opt.device) if opt.device != 'cpu' else 'cpu' model = XLSR(opt.SR_rate) # load pretrained model if opt.model.endswith('.pt') and os.path.exists(opt.model): model.load_state_dict(torch.load(opt.model, map_location=device)) else: model.load_state_dict(torch.load(os.path.join(opt.save_dir, 'best.pt'), map_location=device))
报错信息
Traceback (most recent call last): File "C:\Users\XLSR\test.py", line 110, in <module> model.load_state_dict(torch.load(os.path.join(opt.save_dir, 'best.pt'), map_location=device)) File "C:\Users\AppData\Local\anaconda3\Lib\site-packages\torch\nn\modules\module.py", line 2152, in load_state_dict raise RuntimeError('Error(s) in loading state_dict for {}:\n\t{}'.format( RuntimeError: Error(s) in loading state_dict for XLSR: Missing key(s) in state_dict: "Gblocks.0.conv0.conv2d_block.0.weight", "Gblocks.0.conv0.conv2d_block.0.bias", "Gblocks.0.conv0.conv2d_block.1.weight", "Gblocks.0.conv0.conv2d_block.1.bias", "Gblocks.0.conv0.conv2d_block.2.weight", "Gblocks.0.conv0.conv2d_block.2.bias", "Gblocks.0.conv0.conv2d_block.3.weight", "Gblocks.0.conv0.conv2d_block.3.bias", "Gblocks.1.conv0.conv2d_block.0.weight", "Gblocks.1.conv0.conv2d_block.0.bias", "Gblocks.1.conv0.conv2d_block.1.weight", "Gblocks.1.conv0.conv2d_block.1.bias", "Gblocks.1.conv0.conv2d_block.2.weight", "Gblocks.1.conv0.conv2d_block.2.bias", "Gblocks.1.conv0.conv2d_block.3.weight", "Gblocks.1.conv0.conv2d_block.3.bias", "Gblocks.2.conv0.conv2d_block.0.weight", "Gblocks.2.conv0.conv2d_block.0.bias", "Gblocks.2.conv0.conv2d_block.1.weight", "Gblocks.2.conv0.conv2d_block.1.bias", "Gblocks.2.conv0.conv2d_block.2.weight", "Gblocks.2.conv0.conv2d_block.2.bias", "Gblocks.2.conv0.conv2d_block.3.weight", "Gblocks.2.conv0.conv2d_block.3.bias". Unexpected key(s) in state_dict: "Gblocks.0.conv0.weight", "Gblocks.0.conv0.bias", "Gblocks.1.conv0.weight", "Gblocks.1.conv0.bias", "Gblocks.2.conv0.weight", "Gblocks.2.conv0.bias".
问题原因与解决方法
问题原因
核心是模型结构与预训练权重的结构不匹配:
- 当前本地代码中,
XLSR模型的Gblocks.x.conv0是一个包含多个子层的conv2d_block模块,权重key格式为Gblocks.x.conv0.conv2d_block.x.xxx - 但你加载的
best.pt是用旧版模型结构训练保存的,当时Gblocks.x.conv0是单个卷积层,权重key格式为Gblocks.x.conv0.xxx
这种差异通常是本地代码版本和训练best.pt时的代码版本不一致导致的。
解决方法
方法1:对齐代码版本
找到训练best.pt对应的仓库代码版本,将本地代码切换到该版本,确保模型结构与权重key完全对应,即可直接加载。
方法2:手动修改权重字典的key
如果不想切换代码版本,可加载权重后手动调整key格式,再加载到模型:
# 加载原始权重字典 state_dict = torch.load(os.path.join(opt.save_dir, 'best.pt'), map_location=device) new_state_dict = {} # 遍历并替换key格式 for old_key, value in state_dict.items(): if 'Gblocks.' in old_key and '.conv0.weight' in old_key: new_key = old_key.replace('.conv0.weight', '.conv0.conv2d_block.0.weight') new_state_dict[new_key] = value elif 'Gblocks.' in old_key and '.conv0.bias' in old_key: new_key = old_key.replace('.conv0.bias', '.conv0.conv2d_block.0.bias') new_state_dict[new_key] = value else: new_state_dict[old_key] = value # 加载修改后的权重,strict=False忽略未匹配的其他层(自动初始化) model.load_state_dict(new_state_dict, strict=False)
方法3:重新训练模型
若上述方法都不适用,直接使用当前本地的模型结构重新训练,得到匹配的best.pt后再执行测试。
内容的提问来源于stack exchange,提问作者angel_30
相关产品推荐
相关产品推荐

