如何暂停/终止Optuna研究并从断点或从头恢复未完成试验?
Optuna Study 暂停、终止与恢复方案
当然可以,Optuna支持Study的暂停、终止与恢复,并且能满足你提到的两种恢复需求,下面是具体实现方案:
一、先做基础准备:持久化Study
默认情况下optuna.create_study()会把Study数据存在内存里,进程终止后就会丢失。要实现恢复,必须将Study存储到持久化数据库(比如SQLite、PostgreSQL等):
# 使用SQLite作为持久化存储,数据会存在study.db文件中 study = optuna.create_study( storage="sqlite:///study.db", study_name="my_target_study", # 指定Study名称,后续恢复时需要用到 load_if_exists=True # 如果该名称的Study已存在则直接加载,否则创建新的 )
二、暂停/终止Study的方式
直接终止运行study.optimize()的进程即可,Optuna会自动记录所有已完成的Trials,未完成的Trials会被标记为FAILED或RUNNING(具体取决于终止方式)。
三、两种恢复方案
1. 从头运行未完成的Trials
这是Optuna内置的能力,不需要额外编写逻辑。重新加载已有的Study后再次调用optimize(),Optuna会自动跳过所有已完成的Trials,重新运行之前未完成的那些:
# 加载已存储的Study study = optuna.create_study( storage="sqlite:///study.db", study_name="my_target_study", load_if_exists=True ) # 继续执行优化,Optuna会自动跳过已完成的Trials,重新处理未完成的 study.optimize(objective, n_trials=100) # 可以继续指定需要完成的总Trials数量
注:如果之前的未完成Trial被标记为
RUNNING,Optuna在恢复时会自动将其转为FAILED,然后重新运行。
2. 从最新Checkpoint恢复未完成的Trials
这种需求需要你在objective函数中自行实现Checkpoint的保存与加载逻辑,把每个Trial的中间状态(比如模型权重、训练进度等)保存下来,恢复时加载对应的状态继续执行。
示例代码(以PyTorch训练为例):
import os import torch import optuna def objective(trial): # 定义当前Trial的Checkpoint路径 checkpoint_path = f"./checkpoints/trial_{trial.number}.pth" os.makedirs("./checkpoints", exist_ok=True) # 加载已有的Checkpoint(如果存在) if os.path.exists(checkpoint_path): checkpoint = torch.load(checkpoint_path) start_epoch = checkpoint["epoch"] model = build_model(trial) # 根据Trial参数构建模型 model.load_state_dict(checkpoint["model_state"]) else: start_epoch = 0 model = build_model(trial) # 从上次中断的epoch开始训练 for epoch in range(start_epoch, 100): # 执行训练步骤 loss = train_one_epoch(model) # 向Optuna报告当前epoch的损失 trial.report(loss, epoch) # 保存当前状态到Checkpoint torch.save({ "epoch": epoch + 1, "model_state": model.state_dict(), "trial_params": trial.params }, checkpoint_path) # 如果Optuna建议剪枝该Trial,提前终止并清理Checkpoint if trial.should_prune(): os.remove(checkpoint_path) raise optuna.TrialPruned() # 训练完成后清理Checkpoint(可选) os.remove(checkpoint_path) return loss
恢复时,只需要重新加载Study并调用optimize()即可:Optuna会找到未完成的Trial,调用objective函数时会自动加载对应的Checkpoint,从上次中断的位置继续执行。
关键注意事项
- 必须使用持久化存储,否则进程终止后Study的所有数据都会丢失,无法恢复。
- 从头运行未完成Trials是Optuna原生支持的功能,无需额外代码;而从Checkpoint恢复需要自行实现状态的保存与加载逻辑。
- 如果在分布式环境下使用,要确保所有进程都能访问到同一个存储数据库和Checkpoint文件目录(比如使用共享文件系统)。
内容的提问来源于stack exchange,提问作者gameveloster
相关产品推荐
相关产品推荐

