尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

Ray Tune Callback 与 Metrics 完全指南:自定义回调、日志指标与自动填充字段解析

Ray Tune Callback 与 Metrics 完全指南:自定义回调、日志指标与自动填充字段解析 Ray Tune Callback 与 Metrics 完全指南自定义回调、日志指标与自动填充字段解析【免费下载链接】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 中两个紧密相关的核心能力Callback回调机制与Metrics指标记录体系。你可以在 Function API 与 Class API 训练中记录任意自定义指标并通过继承tune.Callback在训练、恢复、保存、出错等生命周期节点自动触发自定义逻辑同时 Tune 会为每条结果自动填充trial_id、time_total_s、training_iteration等系统指标这些指标可以直接作为停止条件或传递给 Trial Scheduler 与 Search Algorithm。读完本文你将掌握自定义回调的完整钩子清单、两类训练 API 的指标上报方式以及全部自动填充字段的语义与源码依据并能独立编写可投入实际调参任务的回调组件。Ray Tune 回调机制在训练生命周期中插入自定义逻辑Ray Tune 支持在训练过程的不同阶段被自动调用的回调Callback。回调通过参数传递给RunConfig由Tuner接收之后你在回调中实现的方法会在对应时机被自动触发无需手动调用。一个最简回调示例下面的回调会在每收到一条训练结果result时打印其中名为metric的指标值from ray import tune from ray.tune import Callback class MyCallback(Callback): def on_trial_result(self, iteration, trials, trial, result, **info): print(fGot result: {result[metric]}) def train_fn(config): for i in range(10): tune.report({metric: i}) tuner tune.Tuner( train_fn, run_configtune.RunConfig(callbacks[MyCallback()])) tuner.fit()运行后每当train_fn通过tune.report上报一条结果MyCallback.on_trial_result就会被调用一次并打印当前的指标值。RunConfig.callbacks接受一个回调列表因此你可以同时挂载多个回调内部由CallbackList统一聚合分发见 python/ray/tune/callback.py 中的CallbackList实现。完整的回调钩子清单Callback基类python/ray/tune/callback.py为PublicAPI(stabilitybeta)接口定义了以下可覆写的钩子方法钩子方法触发时机关键入参setup(stop, num_samples, total_num_samples, **info)整个实验最开始时调用一次停止条件、采样次数on_step_begin(iteration, trials, **info)每次调参主循环步进开始时循环迭代次数、当前 trialson_step_end(iteration, trials, **info)每次调参主循环步进结束时迭代计数已递增同上on_trial_start(iteration, trials, trial, **info)一个 trial 启动之后对应 trial 对象on_trial_restore(iteration, trials, trial, **info)一个 trial 从 checkpoint 恢复之后对应 trial 对象on_trial_save(iteration, trials, trial, **info)收到一个 trial 保存的 checkpoint 之后对应 trial 对象on_trial_result(iteration, trials, trial, result, **info)收到一个 trial 上报的结果之后result 字典on_trial_complete(iteration, trials, trial, **info)一个 trial 正常完成之后对应 trial 对象on_trial_recover(iteration, trials, trial, **info)一个 trial 出错但被安排重试之后对应 trial 对象on_trial_error(iteration, trials, trial, **info)一个 trial 出错errored之后对应 trial 对象on_checkpoint(iteration, trials, trial, checkpoint, **info)一个 trial 通过 Tune 保存了 checkpoint 之后ray.tune.Checkpoint对象on_experiment_end(trials, **info)实验结束、所有 trial 完结之后全部 trials从源码可以看出on_trial_result、on_trial_complete、on_trial_error这几个钩子在调用前会先通知 Search Algorithm 与 Scheduler而on_trial_recover则不会通知它们见 python/ray/tune/callback.py 中各钩子的 docstring。这意味着在这些回调中读取到的 trial 状态已经反映过调度决策。回调的生命周期与状态持久化Callback还提供了两个用于有状态回调的方法get_state() - Optional[Dict]返回回调当前状态的字典表示默认返回None表示无状态无需持久化。set_state(state: Dict)根据字典恢复回调状态。Tune 会自动周期性地将回调状态 checkpoint 到实验目录CallbackList使用callback-states-{session}.pkl文件见 python/ray/tune/callback.py 中的CKPT_FILE_TMPL。当实验因故障被恢复实验级容错时回调状态会通过set_state恢复。典型实现如下源码 docstring 中的示例from typing import Dict, List, Optional from ray.tune import Callback from ray.tune.experiment import Trial class MyCallback(Callback): def __init__(self): self._trial_ids set() def on_trial_start(self, iteration, trials, trial, **info): self._trial_ids.add(trial.trial_id) def get_state(self) - Optional[Dict]: return {trial_ids: self._trial_ids.copy()} def set_state(self, state: Dict) - Optional[Dict]: self._trial_ids state[trial_ids]仓库测试 python/ray/tune/tests/test_callbacks.py 专门验证了有状态与无状态回调在实验 checkpoint 时的get_state/set_state行为。如果想要更完整的钩子覆盖测试可参考 python/ray/tune/tests/_test_trial_runner_callbacks.py 中的testCallbackSteps与testCallbacksEndToEnd。回调的实战应用示例仓库中的 python/ray/tune/examples/custom_checkpointing_with_callback.py 展示了如何用on_trial_result实现指标改善时保存 checkpoint的自定义逻辑python/ray/tune/examples/logging_example.py 则演示了回调与自定义 trial 命名结合的场景。更完整的钩子签名与说明请查阅 Ray Tune Callbacks API 文档。如何在训练中记录自定义指标在 Function API 和 Class API 两种训练接口中你都可以记录任意自定义指标值。Function API使用tune.reportdef trainable(config): for i in range(num_epochs): ... tune.report({acc: accuracy, metric_foo: random_metric_1, bar: metric_2})从源码看tune.reportpython/ray/tune/trainable/trainable_fn_utils.py实际上调用的是get_session().report(metrics, checkpointcheckpoint)每次调用都会自动递增底层training_iteration计数——注意该迭代的物理含义由你调用report的频率决定并不一定对应一个 epoch。它还支持第二个可选参数checkpoint用于同时上报并持久化一个 checkpoint。Class API在step()中返回指标字典class Trainable(tune.Trainable): def step(self): ... # dont call report here! return dict(accaccuracy, metric_foorandom_metric_1, barmetric_2)在 Class API 中不要在step()里调用tune.report而是直接把指标字典作为step()的返回值。该返回值会被Trainable.train()合并进结果字典合并逻辑见 python/ray/tune/trainable/trainable.py 中的train()方法。重要提示不要用tune.report传输大数据tune.report只适合上报标量指标。不要用它传输模型、数据集等大对象——这会造成巨大的序列化与通信开销显著拖慢整个 Tune 运行。大对象应通过 checkpoint 机制持久化到存储。哪些指标会被 Tune 自动填充Tune 引入了自动填充指标auto-filled metrics的概念。训练过程中除了你自定义上报的值Tune 会自动为每条结果追加下列指标。这些指标全部可以用于停止条件或作为参数传给 Trial Scheduler / Search Algorithm指标名含义config该 trial 的超参数配置date结果被处理时的日期时间字符串done该 trial 是否已结束True/Falseepisodes_total累计 episode 总数针对 RLlib trainableexperiment_id实验唯一 IDexperiment_tag实验唯一标签包含参数值hostnameworker 的主机名iterations_since_restore从 checkpoint 恢复后tune.report被调用的次数node_ipworker 的主机 IPpidworker 进程的进程 IDtime_since_restore从 checkpoint 恢复后经过的秒数time_this_iter_s当前训练迭代的耗时秒即一次 trainable 函数调用或 Class API 中一次_train()调用time_total_s累计总运行时间秒timestamp结果被处理时的时间戳timesteps_since_restore从 checkpoint 恢复后累计的 timestep 数timesteps_total累计 timestep 总数training_iterationtune.report()被调用的次数trial_idtrial 唯一 ID源码层面的自动填充实现这些指标的填充可以追溯到 python/ray/tune/trainable/trainable.py 中的两处实现get_auto_filled_metrics()第 203 行起一次性填充trial_id、date、timestamp、time_this_iter_s、time_total_s、pid、hostname、node_ip、config、time_since_restore、iterations_since_restore、timesteps_since_restore。train()第 288 行起在每次step()之后按需补填done、training_iteration、timesteps_total、episodes_total等字段其中timesteps_total只在结果里提供了timesteps_this_iter增量时才会累计且不会覆盖用户自带的timesteps_total。所有指标的常量定义集中在 python/ray/tune/result.py包括DONE、HOSTNAME、TRIAL_ID、EXPERIMENT_TAG、NODE_IP、PID、EPISODES_TOTAL、TIMESTEPS_TOTAL、TIME_TOTAL_S等。该文件还定义了DEBUG_METRICS不依赖任何迭代就能获取的指标如trial_id、experiment_id、date、timestamp、pid、hostname、node_ip、config以及AUTO_RESULT_KEYS自动填充指标的完整回归清单。在何处查看这些指标以上所有指标都可以在Trial.last_result字典中查看。除此之外它们也会被 Tune 的日志记录器Logger写入实验结果文件result.json与 CSV 进度文件并可通过ResultGrid/ExperimentAnalysis汇总分析。实战用自动填充指标驱动停止条件由于自动填充指标在每次结果中必然存在你可以放心地将其作为停止条件使用例如基于training_iteration限制迭代数、基于time_total_s限制总时长或基于自定义指标如acc设置早停。停止条件的完整用法可参考 tune-stopping.rst调度器与搜索算法的接入方式可参考 tune-run.rst。在自定义回调中你同样可以直接读取result[trial_id]、result[training_iteration]等自动填充字段来实现按迭代数触发动作之类的逻辑无需自行维护计数器。小结本指南覆盖了 Ray Tune 回调与指标体系的两大主线其一通过继承tune.Callback并覆写on_*钩子配合可选的get_state/set_state实现实验级容错即可在 trial 启动、结果上报、checkpoint 保存、出错重试等全生命周期插入自定义逻辑其二无论 Function API 还是 Class API都能上报任意自定义指标同时 Tune 会基于 python/ray/tune/trainable/trainable.py 的自动填充逻辑为每条结果补齐 18 个系统指标这些指标可直接服务于停止条件、调度器与搜索算法。将两者结合——在on_trial_result中读取自动填充指标与自定义指标——即可构建出高度定制化的调参观测与自动化工作流。【免费下载链接】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),仅供参考
返回列表