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

资讯详情

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

Ray Tune 训练 API 完全指南:Function Trainable 与 Class Trainable 实战详解

Ray Tune 训练 API 完全指南:Function Trainable 与 Class Trainable 实战详解 人工智能分布式训练强化学习任务调度模型推理服务【免费下载链接】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 中编写自定义超参数搜索任务核心是回答一个问题如何在 Tune 的调度循环中定义一次训练本文基于 trainable.rst 官方参考文档系统讲解 Tune 提供的两套训练接口——以tune.report()为核心的Function API函数式训练接口与以tune.Trainable子类为核心的Class API类式训练接口。读完本文你将掌握两种 API 的定义方式与适用场景、中间/最终指标的上报方法、reuse_actors复用 Actor 避免昂贵初始化、PlacementGroupFactory精细资源分配以及二者在迭代计数、指标上报、检查点存取等关键行为上的差异对照并得到来自仓库源码python/ray/tune/trainable/trainable.py、python/ray/tune/trainable/function_trainable.py的底层实现佐证。两种训练 API 概览Ray Tune 的训练循环Trial由两种方式驱动Function API编写一个接收config: dict的普通 Python 函数训练函数内部通过调用tune.report()上报指标由 Tune 在 Ray Actor 进程中执行。Class API继承ray.tune.Trainable抽象类实现setup/step/cleanup等生命周期方法由 Tune 将其实例化为 Ray Actor 运行。文档以一个简单目标函数贯穿始终完整示例见 trainable.pydef objective(x, a, b): return a * (x ** 0.5) b下文所有示例都围绕最大化这个目标函数展开便于直观对比两种 API 的写法差异。Function Trainable API函数式训练接口基本用法在训练循环中上报中间指标Function API 的核心约定是训练函数接收一个config字典参数该字典由 Tune 自动填充对应搜索空间search space中为该 Trial 选中的超参数。搜索空间的定义方式参见 key-concepts.rst 中的tune-key-concepts-search-spaces章节。from ray import tune def trainable(config: dict): intermediate_score 0 for x in range(20): intermediate_score objective(x, config[a], config[b]) tune.report({score: intermediate_score}) # 将 score 上报给 Tune tuner tune.Tuner(trainable, param_space{a: 2, b: 4}) results tuner.fit()每个 Trial 被放入一个独立的 Ray Actor 进程多个 Trial 并行运行。每调用一次tune.report()就产生一个训练结果result并自动推进training_iteration计数。tip不要在Trainable类内部调用tune.report()。tune.report()是为函数式训练设计的会话上报接口Class API 应通过step()的返回值上报指标。自定义上报频率仅在结束时上报最终分数默认情况下tune.report()每次调用都会上报但上报频率完全由你控制。例如只在训练结束时上报一次最终分数from ray import tune def trainable(config: dict): final_score 0 for x in range(20): final_score objective(x, config[a], config[b]) tune.report({score: final_score}) # 仅在结束时上报 tuner tune.Tuner(trainable, param_space{a: 2, b: 4}) results tuner.fit()通过函数返回值上报最终指标除了tune.report()Function API 还支持直接从函数返回值中上报一组最终指标def trainable(config: dict): final_score 0 for x in range(20): final_score objective(x, config[a], config[b]) return {score: final_score} # 将 score 上报给 Tune从源码实现看这一机制由 function_trainable.py 中的wrap_function处理它用inspect检查训练函数的签名要求必须包含唯一的config位置参数函数执行完毕后返回值如果是dict会通过get_session().report(output)上报如果是单个数值则以默认指标DEFAULT_METRIC上报最后还会上报一个RESULT_DUPLICATE标记由 TuneController 识别为训练函数已退出并注入doneTrue从而触发 Trial 的停止决策——这就是返回最终指标即结束训练的底层原理。自动填充指标auto-filled metrics无论使用哪种上报方式Ray Tune 都会在用户指标之外自动填充一批系统指标例如iterations_since_restore。常见自动填充指标包括详见 tune-metrics.rst 的tune-autofilled-metrics章节指标含义config该 Trial 的超参数配置training_iterationtune.report()被调用的次数Function APIiterations_since_restore从检查点恢复后tune.report()被调用的次数time_this_iter_s当前训练迭代耗时秒time_total_s累计运行总时长秒timesteps_total/episodes_total累计时间步 / 回合数如 RLlib Trainablepid/hostname/node_ip工作进程 PID、主机名、节点 IPdate/timestamp结果处理时间doneTrial 是否结束trial_id/experiment_id/experiment_tagTrial 与实验标识这些指标都可以直接用作停止条件或传递给 Trial Scheduler / Search Algorithm。其生成逻辑位于 trainable.py 的get_auto_filled_metrics方法中。Function API 的检查点配置函数式训练接口的检查点checkpoint方式与类式接口不同CheckpointConfig中的checkpoint_frequency和checkpoint_at_end不适用于Function API 检查点而是需要手动在训练函数中控制。核心用法保存检查点通过tune.report(metrics, checkpointcheckpoint)上报每个检查点必须与一组指标同时上报以便按指定指标对检查点排序加载检查点通过tune.get_checkpoint()获取该 Trial 最近一次保存的检查点——当 Trial 失败重试、实验恢复或暂停后恢复如 PBT时会被自动填充。更完整的函数式检查点配置与示例见 tune-trial-checkpoints.rst 的tune-function-trainable-checkpointing章节。底层实现上report()与get_checkpoint()定义在 trainable_fn_utils.pyreport()将指标与可选的Checkpoint交给会话层上报并持久化到存储get_checkpoint()返回会话中已加载的最新检查点。同时请注意tune.report()不适合传输大量数据如模型权重、数据集这会显著拖慢 Tune 运行。Class Trainable API类式训练接口子类化tune.Trainable的基本结构Class API 要求继承ray.tune.Trainable核心生命周期方法有三个完整示例见 trainable.pyfrom ray import tune class Trainable(tune.Trainable): def setup(self, config: dict): # config (dict): 一组超参数 self.x 0 self.a config[a] self.b config[b] def step(self): # 会被反复调用 score objective(self.x, self.a, self.b) self.x 1 return {score: score} tuner tune.Tuner( Trainable, run_configtune.RunConfig( # 训练 20 步 stop{training_iteration: 20}, checkpoint_configtune.CheckpointConfig( # 本示例尚未实现检查点见下文 checkpoint_at_endFalse ), ), param_space{a: 2, b: 4}, ) results tuner.fit()作为tune.Trainable的子类Tune 会在独立进程中基于 Ray Actor API 创建Trainable对象三个生命周期方法的分工为setup训练开始时调用一次负责初始化接收 Tune 自动填充的config字典对应搜索空间中为 Trial 选中的超参数step被多次调用每次调用在调优进程中执行一个逻辑训练迭代内部可包含一个或多个真实训练迭代并通过返回值上报指标cleanup训练结束时调用负责释放资源。caution不要在Trainable类内部调用tune.report()。Class API 的指标上报方式是让step()返回指标字典。tipstep()的执行时间需要权衡。经验法则是单次step()应足够长以摊销调度开销通常超过几秒又足够短以周期性上报进度通常不超过几分钟。Class API 的检查点配置类式训练接口支持三种检查点机制手动触发、按频率触发、训练结束时触发。用户通常只需实现Trainable.save_checkpoint和Trainable.load_checkpoint两个方法并在RunConfig的CheckpointConfig中设置checkpoint_frequency、checkpoint_at_end等选项详见 tune-trial-checkpoints.rst 的tune-class-trainable-checkpointing章节。在源码层面trainable.py 对检查点提供了完整支持save()trainable.py调用用户实现的save_checkpoint()将返回的 dict 或路径统一整理后通过_report_class_trainable_checkpoint()持久化到存储返回_TrainingResultrestore()trainable.py根据load_checkpoint()的实现将检查点还原为 dict 或本地目录形式加载并恢复training_iteration、time_total_s等进度指标若step()返回的结果中包含should_checkpoint: True即tune.result.SHOULD_CHECKPOINT则可手动触发检查点——这在抢占式实例spot instance场景下尤为实用。此外Trainable.__init__会在训练进程中将当前工作目录切换到该 Trial 专属的日志目录self.logdir避免同一物理节点上多个 Trial 互相覆盖文件可通过环境变量RAY_CHDIR_TO_TRIAL_DIR0禁用该行为旧的环境变量TUNE_ORIG_WORKING_DIR已弃用见 trainable.py。高级在 Tune 中复用 Actorreuse_actors如果 Trainable 的初始化非常耗时例如加载大型模型可以为每次 Trial 都重新创建进程会带来巨大开销。Tune 提供了reuse_actorsTrue通过TuneConfig传入Tuner在多个超参数组合之间复用同一个 Trainable Python 进程与对象。该特性仅适用于 Class API。复用的前提是你实现了Trainable.reset_config它接收一组新的超参数并完成更新——是否正确更新超参数完全由用户负责。完整示例from time import sleep import ray from ray import tune from ray.tune.tuner import Tuner def expensive_setup(): print(EXPENSIVE SETUP) sleep(1) class QuadraticTrainable(tune.Trainable): def setup(self, config): self.config config expensive_setup() # 使用 reuse_actorsTrue 时只执行一次 self.max_steps 5 self.step_count 0 def step(self): # 从 config 中提取超参数 h1 self.config[hparam1] h2 self.config[hparam2] # 计算简单二次目标函数最优解位于 hparam13 和 hparam25 loss (h1 - 3) ** 2 (h2 - 5) ** 2 metrics {loss: loss} self.step_count 1 if self.step_count self.max_steps: metrics[done] True # 将计算出的 loss 作为指标返回 return metrics def reset_config(self, new_config): # 复用 Actor 时为新的 Trial 更新配置 self.config new_config return True ray.init() tuner_with_reuse Tuner( QuadraticTrainable, param_space{ hparam1: tune.uniform(-10, 10), hparam2: tune.uniform(-10, 10), }, tune_configtune.TuneConfig( num_samples10, max_concurrent_trials1, reuse_actorsTrue, # 启用 Actor 复用避免昂贵的 setup ), run_configray.tune.RunConfig( verbose0, checkpoint_configray.tune.CheckpointConfig(checkpoint_at_endFalse), ), ) tuner_with_reuse.fit()关键点说明setup()中的expensive_setup()只执行一次之后通过reset_config()切换超参数从而显著加速 PBT 等需要频繁切换配置的算法每次复用切换时reset_config()必须返回True表示重置成功若返回FalseTune 将终止该 Actor 并为新 Trial 创建新进程见 trainable.py 中reset()的实现逻辑代码中max_concurrent_trials1与reuse_actorsTrue搭配保证同一时刻只有一个 Trial 占用该 Actor使复用路径清晰可复现。Function API 与 Class API 对比一览文档给出了两种 API 在关键概念上的对照表概念Function APIClass API训练迭代Training Iteration每次调用tune.report递增每次调用Trainable.step递增上报指标Report metricstune.report(metrics)从Trainable.step返回指标保存检查点Saving a checkpointtune.report(..., checkpointcheckpoint)Trainable.save_checkpoint加载检查点Loading a checkpointtune.get_checkpoint()Trainable.load_checkpoint访问配置Accessing config作为参数传入def train_func(config):通过Trainable.setup传入选型建议Function API 代码量少、上手快适合大多数超参搜索场景Class API 结构化更强适合需要精细控制生命周期、手动管理检查点、复用 Actor 或需要实现default_resource_request声明资源需求的场景。高级资源分配让 Trainable 自身分布式化Trainable 自身也可以被分布式执行。如果你的训练函数/类会进一步创建消耗 CPU/GPU 资源的 Ray Actor 或 Task就需要在PlacementGroupFactory中添加更多 bundle为它们预留额外的资源槽位。例如某个 Trainable 类自身需要 1 个 GPU同时还会启动 4 个各占 1 个 GPU 的 Actor则应通过tune.with_resources指定资源强调行为核心写法tuner tune.Tuner( tune.with_resources(my_trainable, tune.PlacementGroupFactory([ {CPU: 1, GPU: 1}, {GPU: 1}, {GPU: 1}, {GPU: 1}, {GPU: 1} ])), run_configRunConfig(namemy_trainable) )要点补充第一个 bundle{CPU: 1, GPU: 1}是 Trainable 自身所在的主 bundle其后每个{GPU: 1}为该 Trainable 启动的子 Actor 预留资源除 CPU/GPU 外还可以指定memory单位字节以及自定义资源类型custom resourcesClass API 还提供default_resource_request类方法允许 Trainable 根据给定配置自动声明每个 Trial 所需的资源从而免去用户在Tuner中手动设置其基类默认返回None见 trainable.py子类可覆写为PlacementGroupFactory对于 Function APItune.with_resources是请求资源的主要方式其实现位于 util.py资源参数既可以是普通资源字典自动转换为PlacementGroupFactory、PlacementGroupFactory实例也可以是接收 config 并返回工厂的可调用对象with_resources会覆盖已有的资源请求使用时需注意。相关 API 索引围绕本节文档Tune 提供了以下配套 API均可在ray命名空间下导入Function API 相关类tune.Checkpoint、tune.TuneContext函数tune.get_checkpoint、tune.get_context、tune.reportTrainableClass API相关构造函数tune.Trainable需要实现的方法Trainable.setup、Trainable.save_checkpoint、Trainable.load_checkpoint、Trainable.step、Trainable.reset_config、Trainable.cleanup、Trainable.default_resource_requestTune Trainable 工具函数数据注入tune.with_parameters将大型参数以引用方式注入训练函数避免序列化开销资源分配tune.with_resources、tune.execution.placement_groups.PlacementGroupFactory、tune.utils.wait_for_gpu调试工具tune.utils.diagnose_serialization、tune.utils.validate_save_restore、tune.utils.util.validate_warmstart总结Function API以tune.report()驱动迭代与指标上报支持上报中间指标、最终指标以及通过返回值上报配置检查点时需手动通过tune.report(..., checkpoint...)与tune.get_checkpoint()完成Class API通过子类化tune.Trainable实现setup/step/cleanup生命周期支持reuse_actors复用 Actor、reset_config热切换超参数、default_resource_request自动声明资源检查点通过save_checkpoint/load_checkpoint管理资源分配借助tune.with_resources与PlacementGroupFactory可以精确表达 Trainable 自身及其衍生 Actor 的 CPU/GPU/内存/自定义资源需求。本文所有代码示例均可直接复制运行需安装 Ray 并具备对应计算资源。更多细节可继续阅读 tune-metrics.rst自动填充指标、tune-trial-checkpoints.rst检查点配置以及 Tune 核心概念搜索空间与训练循环等文档。赞分享人工智能分布式训练强化学习任务调度模型推理服务【免费下载链接】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 实战使用 HyperBandScheduler 对 Trainable 函数做超参搜索与早停Ray Tune 实战使用 HyperBandScheduler 对 Trainable 函数做超参搜索与早停 本文以 Ray Tune 官方示例 hyper人工智能分布式训练强化学习任务调度模型推理服务Trainable Weka Segmentation 插件使用指南Trainable Weka Segmentation 插件使用指南 项目概述 Trainable Weka Segmentation 是 Fiji增强版 IRay RLlib 算法配置完全指南AlgorithmConfig API 详解与实战Ray RLlib 算法配置完全指南AlgorithmConfig API 详解与实战 导读 AlgorithmConfig 是 Ray RLlib位于 p人工智能分布式训练强化学习任务调度模型推理服务上一篇解锁Playnite潜能10个被忽略的高级设置与隐藏功能下一篇Tomcat性能调优终极指南10个实用技巧提升服务器响应速度创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表