Python中模型训练中断后如何恢复已得到的训练结果?
神经网络训练断点续训实现方法
完全可以,你需要用到训练检查点(Checkpoint)机制,核心逻辑是定期将当前训练的所有关键状态序列化存储到硬盘,中断后只需读取检查点文件就能从断点处继续训练,无需从头开始。
需要存储的核心状态
- 模型的所有权重参数
- 优化器的运行状态(比如动量缓存、学习率衰减步数等,避免恢复后优化逻辑断层)
- 当前已完成的训练轮次、步数、当前损失值、学习率调度器状态等自定义超参数
主流框架实现示例
PyTorch 实现
保存检查点
# 每轮训练结束后执行保存逻辑 checkpoint = { "epoch": current_epoch, "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "current_loss": avg_train_loss, "lr_scheduler_state": lr_scheduler.state_dict() # 按需添加自定义参数 } torch.save(checkpoint, "./epoch_{}_checkpoint.pth".format(current_epoch))
恢复训练
# 加载指定路径的检查点文件 checkpoint = torch.load("./epoch_10_checkpoint.pth") model.load_state_dict(checkpoint["model_state_dict"]) optimizer.load_state_dict(checkpoint["optimizer_state_dict"]) lr_scheduler.load_state_dict(checkpoint["lr_scheduler_state"]) start_epoch = checkpoint["epoch"] + 1 # 后续直接调用训练逻辑从start_epoch开始训练即可
TensorFlow/Keras 实现
Keras内置了ModelCheckpoint回调,无需手动写保存逻辑:
训练时配置自动保存
from tensorflow.keras.callbacks import ModelCheckpoint checkpoint_cb = ModelCheckpoint( filepath="./checkpoint_epoch_{epoch:02d}.h5", save_weights_only=False, # 设为False会同时保存优化器状态 save_freq="epoch", # 每轮训练结束自动保存 # 可选配置:save_best_only=True 仅保存验证集效果最优的检查点,节省空间 # 可选配置:max_to_keep=3 仅保留最近3个检查点文件 ) # 训练时传入回调 model.fit(train_dataset, epochs=50, initial_epoch=0, callbacks=[checkpoint_cb])
恢复训练
from tensorflow.keras.models import load_model # 加载已保存的检查点 model = load_model("./checkpoint_epoch_10.h5") # 从第11轮开始继续训练,initial_epoch设为已完成的轮次数 model.fit(train_dataset, epochs=50, initial_epoch=10)
实用提示
- 无需每轮都保存检查点,根据总训练轮次设置每2~5轮保存一次即可,避免保存操作拖慢训练速度
- 仅保存最优检查点的方案更适合训练稳定的场景,可大幅减少硬盘占用
- 检查点恢复的训练效果和连续训练几乎无差异,不会出现精度损失
内容的提问来源于stack exchange,提问作者Thảo M. Hoàng
相关产品推荐
相关产品推荐

