使用Optuna并行计算时如何设置多进程总n_trials?
Optuna分布式优化全局总试验次数固定方案
核心解法:使用全局计数回调控制总次数
Optuna原生的study.optimize的n_trials参数是单进程级别的限制,每个启动的worker都会独立执行对应次数的试验。要实现全局总次数固定,使用官方提供的MaxTrialsCallback回调即可:
- 回调会读取共享存储中整个Study的已完成试验总数,达到阈值后自动停止所有worker的优化任务
- 不受启动的worker进程数量影响,全局总试验数可控
代码示例
import optuna from optuna.study import MaxTrialsCallback # 你的目标函数 def objective(trial): x = trial.suggest_float("x", -1, 1) return x ** 2 if __name__ == "__main__": # 加载共享存储中的分布式Study study = optuna.load_study( study_name="distributed_opt_study", storage="mysql://your_username:your_password@your_host/your_db_name" # 替换为你的共享存储连接串 ) # 启动优化,全局总试验数设为100 study.optimize( objective, n_trials=None, # 单进程不设次数限制,由回调控制全局总数 callbacks=[ MaxTrialsCallback( n_trials=100, # 全局总试验次数 states=(optuna.trial.TrialState.COMPLETE,) # 统计成功完成的试验,如需包含失败/剪枝的试验可自行添加对应状态 ) ] )
低版本兼容方案
如果你使用的Optuna版本低于2.10.0(该版本首次加入MaxTrialsCallback),可使用如下替代方案:
- 所有worker进程的
n_trials参数设置为远大于目标总次数的数值(如目标总次数为100,可设为10000) - 单独运行监控逻辑,定时查询共享存储中Study的已完成试验数,达到目标值后调用
study.stop()接口,所有worker会自动停止优化任务
注意事项
- 必须使用共享存储(MySQL、PostgreSQL、Redis等),本地内存存储、本地文件SQLite存储无法实现多进程试验状态同步,也无法实现全局次数控制
- 由于多个worker可能同时运行试验还未提交结果,最终总试验数可能会比设定值多最多和worker数相等的次数,属于正常的极小误差
内容的提问来源于stack exchange,提问作者Aurelie Navir
相关产品推荐
相关产品推荐

