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

Optuna RandomSampler与TPESampler始终生成相同参数的问题求助

Optuna RandomSampler与TPESampler始终生成相同参数的问题求助

大家好,我最近在使用Optuna进行超参数优化时遇到了一个棘手的问题,希望能得到各位的帮助:

我当前的任务是基于Flax框架做强化学习的超参搜索,因为Flax不支持多进程训练,所以我把训练逻辑封装到了一个多进程的目标函数中。但在运行时发现,无论是使用TPESampler(配置为sampler=optuna.samplers.TPESampler(multivariate=True, n_startup_trials=10, seed=None))还是RandomSampler,初始的多次trial(比如TPESampler的10次启动trial)总是会生成完全相同的参数组合,这完全不符合随机采样的预期。

我已经尝试使用RDBStorage来存储trial数据,避免进程间的状态问题,但情况依然没有改善。以下是我的完整代码:

import multiprocessing
import functools
import optuna
from optuna.storages import RDBStorage
from sqlalchemy import create_engine
import socket
from functools import partial

def multiprocessing_objective_fn(args, trial):
    queue = multiprocessing.Queue()
    p = multiprocessing.Process(target=train_agent, args=(args, trial, queue))
    p.start()
    p.join()
    result = queue.get()
    return result

if __name__ == "__main__":
    from optuna.storages import RDBStorage
    from sqlalchemy import create_engine
    import socket

    args = get_args()
    print(args)

    # Step 2: Create the engine with the specified timeout
    engine = create_engine("sqlite:///optuna_database_250710_test5" + ".db", connect_args={'timeout': 500})
    # Step 3: Use this engine to create the Optuna storage
    storage = RDBStorage("sqlite:///optuna_database_250710_test5" + ".db")

    study = optuna.create_study(
        direction="maximize",
        storage=storage,
        load_if_exists=True,
        study_name=args.exp_name + '__' + args.env_id + '__' + args.run_name + '__' + args.actor,
        pruner=optuna.pruners.HyperbandPruner(),
        sampler=optuna.samplers.TPESampler(multivariate=True, n_startup_trials=10, seed=None)
    )

    objective_fn = functools.partial(multiprocessing_objective_fn, args)
    # objective_fn = functools.partial(train_agent, args)  # 这行是直接调用训练函数的注释版本

    study.optimize(objective_fn, n_trials=args.n_trials, n_jobs=1)

补充说明:我用RDBStorage是为了确保trial的状态能被正确持久化,避免进程间的状态干扰,但即使设置了n_jobs=1,还是会出现每次trial参数完全相同的情况。我怀疑是不是多进程的封装方式和Optuna的采样逻辑有冲突,或者是数据库存储的连接问题?

麻烦各位帮忙分析一下可能的原因,非常感谢!

内容来源于stack exchange

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.08 03:10:06