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

使用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会报错,解决方法有两种:
    1. 保存时改为存dp.module.state_dict()
    2. 加载时处理权重键:
    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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 16:09:01