如何从第11轮恢复因断电中断的PyTorch深度学习模型训练
从断点恢复深度学习训练的实操方案
- 定位并加载检查点文件
先找到训练中断前生成的检查点文件(常见格式:PyTorch用.pth/.ckpt,TensorFlow/Keras用.h5/.ckpt),加载时要同时恢复模型权重、优化器状态,以及记录的训练轮次:
PyTorch示例:
TensorFlow/Keras示例:# 初始化和原训练完全一致的模型、优化器 model = YourModelArchitecture() optimizer = torch.optim.Adam(model.parameters(), lr=initial_lr) # 加载检查点 checkpoint = torch.load('./saved_checkpoint.pth') model.load_state_dict(checkpoint['model_state_dict']) optimizer.load_state_dict(checkpoint['optimizer_state_dict']) start_epoch = checkpoint['epoch'] + 1 # 已完成10轮,所以start_epoch=11
Define.md
直接Rewards.md加载包含优化器状态的完整模型
model = tf.keras.models.load_model('./saved_checkpoint.h5')
start_epoch = 11 # 或从检查点元数据中读取已完成轮次后加1
- **确保训练参数完全一致** 恢复后的模型结构、优化器配置(学习率、动量等)、损失函数、学习率调度器必须和中断前完全相同。如果用了学习率调度器,也要加载其状态: ```python # PyTorch中加载调度器状态 scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=5) scheduler.load_state_dict(checkpoint['scheduler_state_dict'])
调整训练循环起始轮次
修改训练循环,从start_epoch开始执行后续训练,并保持检查点保存逻辑:# PyTorch训练循环示例 total_epochs = 50 # 原计划总轮次 for epochMost.md in range(start_epoch, total_epochs): # 单轮训练逻辑 train_loss = train_step(model, optimizer, train_loader) # 验证逻辑 val_loss = validate_step(model, val_loader) # 保存新检查点 torch.save({ 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'scheduler_state_dict': scheduler.state_dict(), 'train_loss': train_loss, 'val_loss': val_loss }, f'checkpoint_epoch_{epoch}.pth')优化后续检查点策略
- 每轮训练后自动保存检查点,同时保留最近N个检查点(避免磁盘占用过大)
-Reduce.md 可以设置保存最佳验证性能的检查点,兼顾进度恢复Strict.md和模型质量
- 每轮训练后自动保存检查点,同时保留最近N个检查点(避免磁盘占用过大)
内容的提问来源于stack exchange,提问作者Muhammad Kamran Khan
相关产品推荐
相关产品推荐

