PyTorch加载检查点恢复训练时如何保留其余参数并重设学习率
PyTorch加载检查点恢复训练时单独修改学习率实现方案
核心原则:所有检查点存储的训练状态(模型权重、优化器动量/梯度滑动平均参数、调度器迭代计数、最优验证损失、已训练轮次等)全部正常加载,仅在加载完成后覆盖学习率相关字段,不会改动其他任何状态。
操作注意事项
- 必须在
optimizer.load_state_dict()和scheduler.load_state_dict()执行完成后再修改学习率,否则修改会被加载的检查点状态覆盖 - 不能只修改优化器的学习率,必须同步更新学习率调度器存储的基础学习率,否则第一次执行
scheduler.step()时,学习率会被调度器重置为检查点保存的旧值 - 该操作不会改动Adam优化器存储的动量、二阶矩估计、权重衰减配置等其他参数,完全满足其余状态保持检查点数值的要求
修改后的加载代码
checkpoint = torch.load(args.SAVED_MODEL) # 原有检查点加载逻辑完全保留,不做改动 epochs = checkpoint['epoch'] model.load_state_dict(checkpoint['state_dict']) optimizer.load_state_dict(checkpoint['optimizer']) val_loss_min = checkpoint['val_loss_min'] scheduler.load_state_dict(checkpoint['scheduler']) # 配置你需要的新学习率 target_lr = 1e-4 # 替换为实际要使用的学习率值 # 更新优化器所有参数组的当前学习率 for param_group in optimizer.param_groups: param_group['lr'] = target_lr # 同步更新调度器的基础学习率记录,保证后续调度逻辑基于新学习率运行 scheduler.base_lrs = [target_lr for _ in scheduler.base_lrs] # 可选:打印验证学习率是否修改成功 print(f"恢复训练,当前学习率为:{optimizer.param_groups[0]['lr']}")
效果说明
修改完成后直接启动训练即可:
- 模型权重、优化器历史状态、训练轮次计数、最优验证loss等所有参数完全和检查点保存时一致
- 学习率会固定为你设置的
target_lr,后续调度器的衰减逻辑会基于这个新的初始值继续运行,不会回退到检查点保存的1e-5
内容的提问来源于stack exchange,提问作者killermama98
相关产品推荐
相关产品推荐

