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

运行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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.07 06:50:03