Ray Tune 与 AxSearch 集成指南:基于 Ax 贝叶斯优化的超参数调优实战
【免费下载链接】rayRay is an AI compute engine. Ray consists of a core distributed runtime and a set of AI Libraries for accelerating ML workloads.项目地址: https://gitcode.com/gh_mirrors/ra/ray
本篇指南围绕 Ray Tune 官方示例 ax_example.py(由 ax_example.rst 以literalinclude方式嵌入文档)展开,系统讲解如何在 Ray Tune 中接入基于 Ax(Adaptive Experimentation,由 Meta 开源、底层基于 BoTorch/PyTorch 的贝叶斯优化平台)的搜索算法AxSearch,包括参数约束、结果约束、并发限制以及与调度器协同使用。读完本文,你将能独立在 Ray Tune 中搭建一套带约束的贝叶斯超参数调优流程,并理解其底层实现机制。
一、AxSearch 在 Ray Tune 中的定位
Ray Tune 内置了多种超参数优化(HPO)框架的集成,包括 Ax、HyperOpt、Optuna、Nevergrad、BOHB、BayesOpt 等。在 doc/source/tune/examples/index.rst 中,这些示例被归类在 "Hyperparameter optimization frameworks" 一节;在 doc/source/tune/api/suggestion.rst 中,AxSearch与BasicVariantGenerator(随机/网格搜索)等并列作为可选的search_alg。
AxSearch的价值在于:
- 使用 Ax 提供的高效贝叶斯优化策略(底层由 BoTorch 驱动,基于高斯过程建模);
- 原生支持参数约束(如
x1 + x2 <= 2.0)与结果约束(如l2norm <= 1.25),这是多数内置搜索算法不具备的能力; - 可与 Ray Tune 的调度器(如
AsyncHyperBandScheduler)叠加使用,实现在提前停止的同时进行智能采样; - 自动把 Ray Tune 的搜索空间语法转换为 Ax 的搜索空间格式。
从源码看,AxSearch继承自Searcher基类,位于 python/ray/tune/search/ax/ax_search.py,对外统一由 python/ray/tune/search/ax/init.py 导出。
二、环境准备
AxSearch依赖独立的ax-platform库。安装命令:
pip install ax-platform示例源码的 docstring 明确标注了该依赖(见 python/ray/tune/examples/ax_example.py 第 1-6 行)。ax-search模块通过惰性导入方式加载 Ax(try: import ax ... except ImportError),若未安装,实例化AxSearch时会直接抛出断言错误并提示安装命令:
assert ax is not None, """Ax must be installed! You can install AxSearch with the command: `pip install ax-platform`."""同时,源码兼容新旧两代 Ax API:新版使用ax.service.ax_client.ObjectiveProperties定义目标,旧版使用objective_name/minimize参数,_setup_experiment()中通过try/except TypeError自动降级适配。
三、完整示例逐段剖析
以下完整代码即文档所引用的示例 python/ray/tune/examples/ax_example.py,它同时验证了AxSearch可以独立调度器(AsyncHyperBandScheduler)配合使用。
3.1 基准函数:Hartmann6
import time import numpy as np from ray import tune from ray.tune.schedulers import AsyncHyperBandScheduler from ray.tune.search.ax import AxSearch def hartmann6(x): alpha = np.array([1.0, 1.2, 3.0, 3.2]) A = np.array( [ [10, 3, 17, 3.5, 1.7, 8], [0.05, 10, 17, 0.1, 8, 14], [3, 3.5, 1.7, 10, 17, 8], [17, 8, 0.05, 10, 0.1, 14], ] ) P = 10 ** (-4) * np.array( [ [1312, 1696, 5569, 124, 8283, 5886], [2329, 4135, 8307, 3736, 1004, 9991], [2348, 1451, 3522, 2883, 3047, 6650], [4047, 8828, 8732, 5743, 1091, 381], ] ) y = 0.0 for j, alpha_j in enumerate(alpha): t = 0 for k in range(6): t += A[j, k] * ((x[k] - P[j, k]) ** 2) y -= alpha_j * np.exp(-t) return yHartmann6 是超参数调优领域经典的 6 维测试函数,定义域通常为[0, 1]^6,具有多个局部极值,常用于验证优化算法对多峰函数的最小化能力。示例中它作为被优化的目标函数,其全局最小值约为-3.322。
3.2 训练函数:easy_objective
def easy_objective(config): for i in range(config["iterations"]): x = np.array([config.get("x{}".format(i + 1)) for i in range(6)]) tune.report( { "timesteps_total": i, "hartmann6": hartmann6(x), "l2norm": np.sqrt((x**2).sum()), } ) time.sleep(0.02)easy_objective接收 Tune 注入的config字典,把搜索空间中的x1~x6组装成向量,每个迭代步通过tune.report(...)上报三类指标:
timesteps_total:当前迭代步(用于调度器判断进度与提前停止);hartmann6:优化目标(metric),mode="min"表示求最小;l2norm:配置向量的 L2 范数,仅作为结果约束的判定指标,不参与主目标优化。
time.sleep(0.02)用于模拟真实训练中每个 step 的开销,使调度器的提前停止机制更具实际意义。
3.3 配置搜索算法与约束
if __name__ == "__main__": import argparse parser = argparse.ArgumentParser() parser.add_argument( "--smoke-test", action="store_true", help="Finish quickly for testing" ) args, _ = parser.parse_known_args() algo = AxSearch( parameter_constraints=["x1 + x2 <= 2.0"], # Optional. outcome_constraints=["l2norm <= 1.25"], # Optional. ) # Limit to 4 concurrent trials algo = tune.search.ConcurrencyLimiter(algo, max_concurrent=4) scheduler = AsyncHyperBandScheduler()这一小段是示例的核心演示点:
parameter_constraints(参数约束):声明搜索空间内参数之间的线性关系约束,这里x1 + x2 <= 2.0。由于x1、x2本身取值范围是[0, 1],该约束在本例中并不收紧可行域,主要用于演示语法;Ax 支持诸如"x3 >= x4"、"x3 + x4 >= 2"这类表达式。outcome_constraints(结果约束):对上报指标施加约束,"l2norm <= 1.25"表示只接受 L2 范数不超过 1.25 的候选点。Ax 的贝叶斯优化会在采样的同时考虑满足结果约束的概率(即约束贝叶斯优化),这一点在 ax_search.py 的_process_result中也有体现——完成一个 trial 时,除了目标指标,所有outcome_constraints涉及的指标也会一并喂给 Ax(metrics_to_include)。ConcurrencyLimiter(并发限制):Ax 的默认生成策略通常是串行优化的(示例源码在检测到_enforce_sequential_optimization时会提示"Be sure to use a ConcurrencyLimiter"),通过tune.search.ConcurrencyLimiter(algo, max_concurrent=4)包装后,最多同时运行 4 个 trial,避免因并行采样过多而降低贝叶斯模型质量,同时仍能充分利用多核/多机资源。AsyncHyperBandScheduler:异步超带调度器,根据timesteps_total的进度提前终止表现不佳的 trial,与搜索算法正交组合,示例专门验证了"带独立调度器"场景的可用性。
3.4 Tuner 组装与运行
tuner = tune.Tuner( easy_objective, run_config=tune.RunConfig( name="ax", stop={"timesteps_total": 100}, ), tune_config=tune.TuneConfig( metric="hartmann6", # provided in the 'easy_objective' function mode="min", search_alg=algo, scheduler=scheduler, num_samples=10 if args.smoke_test else 50, ), param_space={ "iterations": 100, "x1": tune.uniform(0.0, 1.0), "x2": tune.uniform(0.0, 1.0), "x3": tune.uniform(0.0, 1.0), "x4": tune.uniform(0.0, 1.0), "x5": tune.uniform(0.0, 1.0), "x6": tune.uniform(0.0, 1.0), }, ) results = tuner.fit() print("Best hyperparameters found were: ", results.get_best_result().config)关键点:
metric="hartmann6"必须与easy_objective中tune.report的键一致,mode="min"明确优化方向;AxSearch构造时未显式传metric/mode,说明二者可以通过TuneConfig注入(内部由set_search_properties完成)。param_space中x1~x6均为tune.uniform(0.0, 1.0)连续量,iterations=100是固定超参(不参与搜索)。AxSearch会自动将 Tune 语法转换为 Ax 的 range 参数。num_samples控制总 trial 数:正常跑 50 个,--smoke-test时仅 10 个用于快速验证。- 结果通过
results.get_best_result().config直接打印最优超参数配置。
运行方式:
# 完整运行 50 个 trial python python/ray/tune/examples/ax_example.py # 冒烟测试,仅 10 个 trial,快速验证流程 python python/ray/tune/examples/ax_example.py --smoke-test四、AxSearch 核心 API 与参数详解
根据 ax_search.py 的类文档与构造实现,AxSearch的完整签名如下:
AxSearch( space=None, # 手动指定的 Ax 搜索空间,字典或列表形式 metric=None, # 优化指标名,须与 tune.report 中的键一致 mode=None, # "min" 或 "max" points_to_evaluate=None, # 初始建议点,列表[dict] parameter_constraints=None,# 参数约束,如 "x3 >= x4" outcome_constraints=None, # 结果约束,如 "m1 <= 3" ax_client=None, # 已初始化的 AxClient 实例 **ax_kwargs, # 传递给 AxClient 的其他参数(如 random_seed) )各参数说明:
| 参数 | 类型 | 含义与说明 |
|---|---|---|
space | dict/list[dict] | Ax 格式搜索空间。若为 Tune 格式字典,会被自动转换;若不传,则必须通过Tuner(param_space=...)或已有ax_client提供 |
metric | str | 目标指标名,必须出现在tune.report的结果字典中;若为None但指定了mode,默认使用ray.tune.result.DEFAULT_METRIC |
mode | "min"/"max" | 优化方向,默认"max"。注意示例中显式设为"min" |
points_to_evaluate | list[dict] | 已有先验好配置,按顺序最先运行,帮助算法预热 |
parameter_constraints | list[str] | 线性参数约束表达式,如"x3 >= x4"、"x3 + x4 >= 2" |
outcome_constraints | list[str] | 形如"metric_name >= bound"的结果约束,如"m1 <= 3" |
ax_client | AxClient | 复用一个已创建好实验的AxClient;此时不得再传space、metric、parameter_constraints、outcome_constraints(源码会显式抛ValueError校验) |
**ax_kwargs | - | 透传给内部AxClient(**kwargs),例如random_seed;当显式传入ax_client时被忽略 |
4.1 方式一:自动转换 Tune 搜索空间(推荐)
from ray import tune from ray.tune.search.ax import AxSearch config = { "x1": tune.uniform(0.0, 1.0), "x2": tune.uniform(0.0, 1.0), } def easy_objective(config): for i in range(100): intermediate_result = config["x1"] + config["x2"] * i tune.report({"score": intermediate_result}) ax_search = AxSearch() tuner = tune.Tuner( easy_objective, tune_config=tune.TuneConfig( search_alg=ax_search, metric="score", mode="max", ), param_space=config, ) tuner.fit()AxSearch会调用静态方法convert_search_space将 Tune 采样器转换为 Ax 参数定义。转换规则(见 ax_search.py 的convert_search_space):
| Tune 采样器 | Ax 参数类型 | 说明 |
|---|---|---|
tune.uniform(a, b)(Float) | {"type": "range", "bounds": [a, b], "value_type": "float", "log_scale": False} | 连续均匀 |
tune.loguniform(a, b)(Float) | {"type": "range", "bounds": [a, b], "value_type": "float", "log_scale": True} | 对数均匀 |
tune.uniform(a, b)(Integer) | {"type": "range", "bounds": [a, b-1], "value_type": "int", ...} | 整型均匀,注意上界减 1 |
tune.loguniform(a, b)(Integer) | 同上,log_scale: True | 整型对数均匀 |
tune.choice([...])(Categorical) | {"type": "choice", "values": categories} | 类别型 |
| 嵌套 dict / list 中的固定值 | {"type": "fixed", "value": val} | 固定参数 |
需要注意的限制:
- 不支持
grid_search:转换时若检测到grid_vars会直接抛ValueError("Grid search parameters cannot be automatically converted to an Ax search space."); - 不支持量化采样器:
tune.quniform等带Quantized包装的采样器会打印警告并丢弃量化("AxSearch does not support quantization. Dropped quantization."); - 嵌套 dict/list 的参数名以
/连接,例如"a/b/0"。
4.2 方式二:手动传递 Ax 格式搜索空间
from ray import tune from ray.tune.search.ax import AxSearch parameters = [ {"name": "x1", "type": "range", "bounds": [0.0, 1.0]}, {"name": "x2", "type": "range", "bounds": [0.0, 1.0]}, ] def easy_objective(config): for i in range(100): intermediate_result = config["x1"] + config["x2"] * i tune.report({"score": intermediate_result}) ax_search = AxSearch(space=parameters, metric="score", mode="max") tuner = tune.Tuner( easy_objective, tune_config=tune.TuneConfig(search_alg=ax_search), ) tuner.fit()Ax 参数字典必含字段:name(参数名)、type("range"/"fixed"/"choice")、range 类型需bounds(下界在前),choice 类型需values,fixed 类型需单个value。该方式便于复用已有的 Ax 实验定义,或精确控制log_scale等细节。
4.3 方式三:复用已有 AxClient(高级)
当需要复用AxClient(例如跨实验共享随机种子、或接入已有 Ax 实验)时:
from ax.service.ax_client import AxClient, ObjectiveProperties from ray.tune.search.ax import AxSearch client = AxClient(random_seed=4321) client.create_experiment( parameters=converted_config, objectives={"_metric": ObjectiveProperties(minimize=False)}, ) searcher = AxSearch(ax_client=client)这正是 python/ray/tune/tests/test_searchers.py 中testAxManualSetup的用法。此时AxSearch不再自行创建实验,而是直接向已有实验追加 trial,且构造参数中不允许再携带任何实验定义信息(源码中的冲突校验逻辑见_setup_experiment)。
五、约束机制详解
约束是AxSearch区别于多数内置搜索算法的杀手锏,示例中同时演示了两种:
5.1 参数约束(parameter_constraints)
对搜索空间内参数施加线性不等式,例如:
AxSearch(parameter_constraints=["x1 + x2 <= 2.0"])Ax 支持任意参数的线性组合表达式,包括"x3 >= x4"、"x3 + x4 >= 2"等形式。贝叶斯优化在推荐下一组参数时会把该约束纳入采样过程,避免无效探索。
5.2 结果约束(outcome_constraints)
对训练过程中上报的指标施加边界:
AxSearch(outcome_constraints=["l2norm <= 1.25"])其形式为"指标名 比较符 边界"(如"m1 <= 3")。从 ax_search.py 的_process_result可以看到,trial 完成时 Ax 需要同时收到目标指标与所有结果约束指标的值:
metrics_to_include = [self._metric] + [ oc.metric.name for oc in self._ax.experiment.optimization_config.outcome_constraints ]即 Ax 会用带约束的高斯过程模型同时建模目标与约束的可行性,推荐"大概率满足约束且目标优秀"的候选点。
六、源码级运行原理
6.1 一次 trial 的生命周期
AxSearch实现Searcher接口的两个核心方法:
suggest(trial_id):若无points_to_evaluate剩余,则调用self._ax.get_next_trial()获取 Ax 推荐的新参数,并将 Tune 的trial_id映射到 Ax 的trial_index(保存在_live_trial_mapping)。若 Ax 因并行上限(MaxParallelismReachedException)或数据不足(DataRequiredError)无法给出新点,则返回None让 Tune 暂停生成。有初始建议点时,则从points_to_evaluate中弹出并通过attach_trial(config)直接附加。on_trial_complete(trial_id, result, error):将结果交给_process_result,其中若发现metric为 NaN/Inf,会调用ax.abandon_trial()放弃该 trial 而不是上报非法值;合法结果则通过ax.complete_trial(trial_index, raw_data=metric_dict)回填,驱动贝叶斯模型更新。
返回值还需经过unflatten_list_dict还原为嵌套配置;当搜索空间含"固定值与可调参数混排的列表"(如[1, tune.uniform(2, 3), 4])时,会先对键排序再反扁平化,避免键序错乱。
6.2 检查点与恢复
AxSearch通过save(checkpoint_path)/restore(checkpoint_path)用cloudpickle序列化整个实例状态,实现 Tune 的断点续训与故障恢复。test_searchers.py中的check_searcher_checkpoint_errors_scope专门校验搜索算法检查点化后不出现序列化错误。
6.3 与调度器的协同
示例使用AsyncHyperBandScheduler验证了搜索算法与调度器可叠加:调度器负责"何时提前终止表现差的 trial",搜索算法负责"下一个 trial 采哪里"。Ray Tune 的ConcurrencyLimiter则进一步协调两者——控制同时在跑的 trial 数,确保 Ax 的串行生成策略不被并发打乱。test_convergence.py与test_tune_restore_warm_start.py中同样覆盖了AxSearch与ConcurrencyLimiter组合(如ConcurrencyLimiter(AxSearch(...), max_concurrent=10))及 warm-start 场景。
七、验证与测试
仓库为AxSearch提供了多维度测试,可作为集成正确性的参考:
- python/ray/tune/tests/test_searchers.py:
testAx(自动转换搜索空间 +ConcurrencyLimiter+ 16 个样本,确保 Ax 真正拟合了代理模型)、testAxManualSetup(手动AxClient创建实验,并覆盖混合列表参数); - python/ray/tune/tests/test_convergence.py:收敛性验证(当前
ax warm start用例被标记为跳过,说明 warm-start 依赖的 Ax 版本升级后曾出现问题,集成时建议以当前 Ax 版本实测为准); - python/ray/tune/tests/test_tune_restore_warm_start.py:利用
AxSearch.convert_search_space构造空间、组装AxClient并指定 Ax 生成策略(GenerationStrategy)的恢复/热启动测试,代码中同时兼容 Ax 1.0+(ax.adapter.registry.Generators)与 Ax 0.x(ax.modelbridge.registry.Models)两代 API。
八、实践要点与注意事项
- 务必安装
ax-platform,否则AxSearch构造即失败; metric/mode二选一设置:可在AxSearch(...)传入,也可在TuneConfig传入(set_search_properties会回填);若构造时只给mode不给metric,会退回DEFAULT_METRIC;- 不要对 AxSearch 使用
grid_search,会自动转换失败;量化采样器会被静默降级为普通均匀采样; - 并行度要用
ConcurrencyLimiter控制:Ax 的串行优化策略在未限制并发时可能打印告警,推荐max_concurrent取 4~10 量级; - NaN/Inf 指标会被自动放弃(
abandon_trial),无需在训练函数中额外处理; - 约束语法是字符串表达式:参数约束面向参数名,结果约束面向
tune.report中出现的指标名,写错名称会在 Ax 建实验时报错; - 版本兼容:
ax_search.py对 Ax 新旧 API(ObjectivePropertiesvsobjective_name、GeneratorsvsModels)做了双轨适配,集成时尽量使用较新的ax-platform以获得完整功能。
九、小结
本文以 Ray Tune 官方 Ax 示例为骨架,完整还原了AxSearch的配置、约束、并发控制、调度器组合与运行方式,并结合 python/ray/tune/search/ax/ax_search.py 的实现剖析了搜索空间转换、trial 生命周期、NaN 处理与检查点机制。相比普通随机/网格搜索,AxSearch的贝叶斯优化能在更少的 trial 数内逼近全局最优,尤其适合评估代价高昂的深度学习训练场景;其参数约束与结果约束能力,则为工程上常见的"资源/质量双约束"调优提供了开箱即用的方案。读者可将示例中的hartmann6替换为真实模型训练目标,直接落地到自己的调优流水线中。
【免费下载链接】rayRay is an AI compute engine. Ray consists of a core distributed runtime and a set of AI Libraries for accelerating ML workloads.项目地址: https://gitcode.com/gh_mirrors/ra/ray
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考