PyTorch模型训练后测试模式下查找最佳准确率对应epoch的方法
如何从PyTorch最优模型checkpoint中获取最佳准确率对应的训练轮次
需求背景
我需要确定已保存的最优模型取得最佳准确率对应的epoch,以此判断合理的训练轮数:如果10个epoch就能达到最优效果就无需训练30轮,若10轮尚未收敛则可以继续增加训练轮数。
现有代码与运行结果
当前加载checkpoint的代码如下:
best_acc = 0 # optionally resume from a checkpoint if args.resume: if os.path.isfile(args.resume): print("=> loading checkpoint '{}'".format(args.resume)) checkpoint = torch.load(args.resume) args.start_epoch = checkpoint['epoch'] print(" start epoch: ", args.start_epoch) best_acc = checkpoint['best_prec1'] print("best acc:" , best_acc) model.load_state_dict(checkpoint['state_dict']) print("=> loaded checkpoint '{}' (epoch {})" .format(args.resume, checkpoint['epoch'])) else: print("=> no checkpoint found at '{}'".format(args.resume))
运行后输出如下:
=> loading checkpoint 'runs/both_attn_30e_CA/model_best.pth.tar' start epoch: 6 best acc: 1.4967916199999998 => loaded checkpoint 'runs/both_attn_30e_CA/model_best.pth.tar' (epoch 6)
我训练30个epoch后,使用如下命令在测试模式下运行代码:
python main.py --test --use_fc --resume runs/both_attn_30e_CA/model_best.pth.tar
问题说明
运行结果显示加载的checkpoint的epoch为6,该值同时被赋值给start_epoch。我曾尝试在谷歌搜索checkpoint['best_prec1']相关内容,仅返回2条结果,没有教程说明如何获取best_prec1对应的epoch值,我的训练好的模型保存在runs/both_attn_30e_CA/model_best.pth.tar路径下。
解答
你当前加载的model_best.pth.tar本身就是训练过程中准确率最高的轮次对应的checkpoint,你输出的checkpoint['epoch']的值为6,就说明这个最优准确率就是在第6个epoch取得的,不需要额外查找其他字段对应。
常规PyTorch训练保存最优模型的逻辑为:
- 每个epoch结束后计算验证集准确率,当本轮准确率高于历史最高的
best_prec1时,就会把当前的epoch编号、模型权重、当前更新后的best_prec1等信息打包保存为model_best.pth.tar - 该文件仅保存最优轮次的相关信息,所以读取到的
epoch字段就是最优准确率对应的训练轮次
如果需要后续更方便地判断模型收敛速度,可以修改训练阶段的逻辑,额外把每个epoch的编号、验证集准确率打印到日志文件或者存储为单独的json文件,后续直接查阅该文件就可以直观看到准确率随epoch变化的趋势,无需从checkpoint反向查询。
内容的提问来源于stack exchange,提问作者Mona Jalal
相关产品推荐
相关产品推荐

