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
相关产品推荐
相关产品推荐

