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

PyTorch断点续训加载state_dict时XceptionHourglass缺key报错如何解决

报错核心原因
  • 你保存的checkpoint是自定义结构化字典,包含model/optimizer/epoch/loss四个字段,但加载时直接将整个字典传入load_state_dict方法,没有取出对应字段的权重数据,导致模型在字典里找不到匹配的参数键。
  • 如果你的预训练权重是通过nn.DataParallel包装后的模型直接保存的,权重键会自带module.前缀,和当前直接定义的XceptionHourglass模型的参数键不匹配,也会触发该报错。
修复步骤

1. 修正加载逻辑,取出对应字段

将原加载代码替换为以下内容,首先取出字典内的对应字段再加载:

checkpoint = torch.load('imaterialist2020-pretrain-models/maskmodel_160.model_ep4_tsave')
# 加载模型权重
model.load_state_dict(checkpoint['model'])
# 加载优化器状态
optimizer.load_state_dict(checkpoint['optimizer'])
start_epoch = checkpoint['epoch']
loss = checkpoint['loss']

2. 调整训练起始epoch

原循环从0开始迭代,会覆盖你加载的断点进度,修改循环起始值:

# 从断点的下一个epoch开始训练,总训练轮数保持num_epochs不变
for epoch in range(start_epoch + 1, start_epoch + num_epochs):

3. 可选:处理DataParallel前缀

如果完成前两步后仍报参数键缺失错误,说明预训练权重是nn.DataParallel包装后保存的,需要手动去掉参数键的module.前缀:

checkpoint = torch.load('imaterialist2020-pretrain-models/maskmodel_160.model_ep4_tsave')
raw_state_dict = checkpoint['model']
processed_state_dict = {}
for k, v in raw_state_dict.items():
    # 移除module.前缀
    if k.startswith('module.'):
        processed_state_dict[k[7:]] = v
    else:
        processed_state_dict[k] = v
model.load_state_dict(processed_state_dict)
# 优化器加载逻辑不变
optimizer.load_state_dict(checkpoint['optimizer'])
start_epoch = checkpoint['epoch']
loss = checkpoint['loss']
额外说明

你当前的模型保存逻辑是直接存储原始XceptionHourglass实例的state_dict,后续自己保存的断点再加载时不需要重复处理前缀,仅旧的DataParallel直接保存的权重需要走第三步的处理逻辑。

内容的提问来源于stack exchange,提问作者Entropie_13

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 18:39:03