使用torch.save与torch.load继续PyTorch训练出现key error报错问题
报错核心原因
你当前的保存逻辑仅存储了模型的权重字典:
torch.save(model.state_dict(), save_path)
这个操作输出的文件是纯OrderedDict结构,仅包含模型层的参数值,没有你需要的优化器状态、当前训练轮次、损失值等字段,所以你调用checkpoint['epoch']、checkpoint['loss']这类操作必然触发KeyError。
两种续跑方案
方案1:直接使用现有已保存的文件(仅恢复模型权重)
如果不需要完全还原训练状态,仅要基于已训练的权重继续训练,可以直接加载现有文件,不需要修改保存逻辑,对应代码调整如下:
# 加载部分保留现有写法,注释掉优化器、epoch、loss的加载逻辑即可 checkpoint = torch.load('imaterialist2020-pretrain-models/maskmodel_160.model_ep17') model.load_state_dict(checkpoint)
注意:这种方案下优化器、学习率调度器会从头初始化,训练轮次也会从0开始,属于「预训练权重微调」,不是完整的断点续跑。
方案2:编写checkpoint存储逻辑实现完整续跑
如果要完全还原训练状态(包括优化器动量、学习率进度、训练轮次),需要调整保存逻辑,将所有需要的状态打包成字典存储:
第一步:修改保存逻辑
建议将每step保存改为每epoch结束后保存,避免产生大量冗余文件,修改如下:
# 把循环内step级的torch.save注释,移到epoch循环末尾 for epoch in range(num_epochs): # ... 原有训练逻辑不变 ... # epoch结束后存储完整checkpoint checkpoint = { 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'scheduler_state_dict': lr_scheduler.state_dict() } save_name = f"maskmodel_{attr_image_size[0]}_ep{epoch}_tsave.pt" torch.save(checkpoint, os.path.join(MODEL_FILE_DIR, save_name))
第二步:修改加载逻辑
增加存在性判断,不存在checkpoint则从零开始训练:
start_epoch = 0 checkpoint_path = '你的checkpoint文件路径' if os.path.exists(checkpoint_path): checkpoint = torch.load(checkpoint_path) # 加载模型权重 model.load_state_dict(checkpoint['model_state_dict']) # 加载优化器状态 optimizer.load_state_dict(checkpoint['optimizer_state_dict']) # 加载学习率调度器状态 lr_scheduler.load_state_dict(checkpoint['scheduler_state_dict']) # 设置起始训练轮次 start_epoch = checkpoint['epoch'] + 1 # 训练循环改为从start_epoch开始 for epoch in range(start_epoch, num_epochs): # ... 原有训练逻辑不变 ...
注意事项
- 你使用了
torch.nn.DataParallel包装模型,如果保存时直接存dp.state_dict(),权重键会自带module.前缀,直接加载到原始model会报错,解决方法有两种:- 保存时改为存
dp.module.state_dict() - 加载时处理权重键:
new_state_dict = {k.replace('module.', ''): v for k, v in checkpoint['model_state_dict'].items()} model.load_state_dict(new_state_dict) - 保存时改为存
- 交叉熵损失函数是无状态的,不需要存储到checkpoint中,每次初始化即可。
内容的提问来源于stack exchange,提问作者Entropie_13
相关产品推荐
相关产品推荐

