You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Optuna中断/恢复超参数搜索结果无法复现问题咨询

Optuna中断恢复后与同种子不间断运行结果不一致的解决方案

问题概述

针对大参数规模的ML模型,Optuna的中断恢复优化功能实用性很强,但固定随机种子时,中断恢复后的研究结果与单次不间断运行的结果出现差异:

  • 预期:三种场景(单次n_trials=x、分5次恢复累计x次、键盘中断5次累计x次)的优化结果完全一致
  • 实际:首次中断前结果匹配,中断恢复后结果偏离

核心疑问

目标函数无随机性时,能否实现中断恢复后与同种子不间断运行结果完全一致?

复现代码

import optuna
import logging
import sys
import numpy as np

def objective(trial):
    x = trial.suggest_float("x", -10, 10)
    return (x - 4) ** 2

def set_study(db_name, 
                study_name, 
                seed, 
                direction="minimize"):
    '''创建可中断恢复的Optuna研究'''
    optuna.logging.get_logger("optuna").addHandler(logging.StreamHandler(sys.stdout))
    sampler = optuna.samplers.TPES皆保持 Fast计划 lesser Energyial**al/group�VM
awt一组到什么程度?不对,原代码保留:
    sampler = optuna.samplers.TPESampler(seed=seed, n_startup_trials=0)
    storage_name = f"sqlite:///{db_name}.db"
    storage = optuna.storages.RDBStorage(storage_name, heartbeat_interval=1)

    study = optuna.create_study(storage=storage, 
                                study_name=study_name, 
                                sampler=sampler, 
                                direction=direction, 
                                load_if_exists=True)
    return study

study = set_study('optuna_test', 'optuna_test_study', 1)

try:
    # 按CTRL+C触发中断
    study.optimize(objective, n_trials=100)  
except KeyboardInterrupt:
    pass

# 输出结果
df = study.trials_dataframe(attrs=("number", "value", "params", "state"))
print(df)
print("Best params: ", study.best_params)
print("Best value: ", study.best_value)

# 可视化
fig = optuna.visualization.plot_optimization_history(study)
fig.show()

解决方案:持久化采样器内部状态

可以实现完全一致的结果,核心是保存并恢复采样器的内部状态。默认情况下,Optuna不会将TPESampler的内部状态(如随机数生成器状态、历史统计信息)写入RDB存储,中断恢复时采样器重新初始化,导致后续采样偏离。

具体实现步骤

  1. 修改研究初始化逻辑:加载研究时,从用户属性中恢复采样器状态(如果存在)
  2. 添加状态保存回调:在每次试验结束后,将采样器状态保存到研究的用户属性中

修改后的代码:

import optuna
import logging
import sys
import numpy as np

def objective(trial):
    x = trial.suggest_float("x", -10, 10)
    return (x - 4) ** 2

def save_sampler_state(study, trial):
    '''回调函数:每次试验后保存采样器状态'''
    study.set_user_attr("sampler_state", study.sampler.__getstate__())

def set_study(db_name, 
                study_name, 
                seed, 
                direction="minimize"):
    optuna.logging.get_logger("optuna").addHandler(logging.StreamHandler(sys.stdout))
    storage_name = f"sqlite:///{db_name}.db"
    storage = optuna.storages.RDBStorage(storage_name, heartbeat_interval=1)

    # 创建/加载研究
    study = optuna.create_study(storage=storage, 
                                study_name=study_name, 
                                direction=direction, 
                                load_if_exists=True)
    
    # 初始化或恢复采样器
    if "sampler_state" in study.user_attrs:
        sampler = optuna.samplers.TPESampler(seed=seed, n_startup_trials=0)
        sampler.__setstate__(study.user_attrs["sampler_state"])
    else:
        sampler = optuna.samplers.TPESampler(seed=seed, n_startup_trials=0)
    
    study.sampler = sampler
    return study

study = set_study('optuna_test', 'optuna_test_study', 1)

try:
    # 添加状态保存回调
    study.optimize(objective, n_trials=100, callbacks=[save_sampler_state])  
except KeyboardInterrupt:
    # 中断时手动保存最后状态
    save_sampler_state(study, None)
    pass

# 输出结果
df = study.trials_dataframe(attrs=("number", "value", "params", "state"))
print(df)
print("Best params: ", study.best_params)
print("Best value: ", study.best_value)

fig = optuna.visualization.plot_optimization_history(study)
fig.show()

关键注意事项

  • 确保每次恢复研究时使用完全相同的种子、存储配置和研究名称
  • 中断时手动触发状态保存,避免最后一次试验的状态丢失
  • 该方法适 landAmb这ind外的 production Goesialsimple r unite副本温馨Reference 不对,重新说:该方法适用于所有实现了__getstate__和__setstate__序列化接口的Optuna采样器

原理说明

Optuna的采样器(如TPESampler)的采样逻辑不仅依赖随机种子,还依赖运行过程中积累的内部状态(比如用于调整采样分布的历史观测数据、随机数生成器的当前状态)。如果不持久化这些状态,恢复时采样器会从头初始化,即使种子相同,也会因为内部状态的差异生成不同的采样结果。通过序列化保存采样器状态,可以完全还原中断前的采样上下文,保证后续采样与不间断运行完全一致。

内容的提问来源于stack exchange,提问作者Maximilian Wirth

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.23 03:54:28