资讯动态

TensorFlow Model Garden Vision 自定义训练启动器实战:扩展 base_trainer 与 train.py 训练驱动器的完整方法

发布时间:2026/9/7 2:13:15 来源:尧图企业网站定制
TensorFlow Model Garden Vision 自定义训练启动器实战扩展 base_trainer 与 train.py 训练驱动器的完整方法【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/models本文基于 TensorFlow Model GardenTFM仓库的文档 customize_training_launcher.md 展开讲解 vision 模块两条核心扩展路径如何继承 base_trainer.Trainer 编写自定义 Trainer 以接管训练/评估循环的关键环节以及如何从标准 train.py 训练驱动器分叉出自己的启动脚本。读完本文你能够针对特定任务修改训练循环、输入获取与指标汇总逻辑并完整复制一份可运行的自定义训练驱动器将其接入 gin 配置与注册表体系。一、先看清 TFM Vision 的启动链路Trainer 从哪里来在动手定制之前有必要先理解标准训练链路中各个组件的关系这决定了“该改哪一层”official/vision/train.py训练驱动器 / 入口脚本 └─ task_factory.get_task(...) # 由注册表按配置创建 Task └─ train_lib.run_experiment(...) # 创建 Trainer 并管理 train/eval 流程 └─ base_trainer.Trainer # 标准训练器实现 Orbit 接口 └─ orbit.Controller # 执行 train / evaluate / train_and_evaluatetrain.py 是 vision 模块的启动脚本负责解析 gin 配置、构造分布式策略tf.distributeStrategy、创建 Task并把一切交给 train_lib.run_experiment。run_experiment内部创建OrbitExperimentRunner它会在strategy.scope()内按需实例化base_trainer.Trainer并通过orbit.Controller按mode分发到不同执行路径见 train_lib.py。Trainer是“模型无关、任务无关”的通用训练器其内部实现依赖 Orbit 的StandardTrainer/StandardEvaluator接口见 base_trainer.py 顶部注释因此在 GPU/TPU/单机之间可以互换。文档中给出的定制动机是当某个具体用例无法用现有的训练函数直接处理时例如自定义的训练循环、特殊的指标聚合、额外的日志或 checkpoint 行为替换或修改 base trainer 的行为能获得对训练过程更大的控制力。二、自定义 Trainer动机与可覆写钩子2.1 动机文档“Motivation”一节明确指出定制 Trainer 的一个典型原因是要替换或修改现有 base trainer 的行为尤其当特定问题需要一种现有训练函数无法轻易处理的独特方法时。定制 Trainer 可以带来对训练流程更细粒度的灵活性与控制从而在具体任务上取得更好的效果。2.2 base Trainer 暴露了哪些可覆写的“钩子”从 base_trainer.py 的源码看Trainer类继承自_AsyncTrainer后者再继承orbit.StandardTrainer与orbit.StandardEvaluator见 base_trainer.py并针对同步/异步两种分布式策略做了统一封装init_async、join、create_train_loop_fn等见 base_trainer.py。文档建议在子类中覆写以下方法钩子方法源码位置职责典型定制场景train_step(iterator)base_trainer.py单步训练取 batch → 调task.train_step→ 更新train_loss、global_step自定义训练循环、条件化编译XLA/jit_compileeval_step(iterator)base_trainer.py单步评估调task.validation_step→ 聚合各副本输出自定义验证循环、额外 passthrough 日志next_train_inputs(iterator)base_trainer.py从迭代器获取下一个训练输入默认return next(iterator)控制送入模型的原始输入如截断、重组 batchnext_eval_inputs(iterator)base_trainer.py获取评估输入并附带passthrough_logs保留在 host 侧不进加速器在评估中携带无法上 GPU/TPU 的额外信息train_loop_end()base_trainer.py每个训练循环结束时汇总 metrics、训练 loss、学习率并重置状态记录多优化器学习率、追加自定义日志eval_end(aggregated_logs)base_trainer.py评估结束汇总 validation metrics/loss、reduce_aggregated_logs、调用checkpoint_exporter.maybe_export_checkpoint、处理 EMA 权重交换自定义最优 checkpoint 导出策略、指标后处理__init__的构造参数同样值得了解见 base_trainer.pyconfigExperimentConfig、task、model、optimizer以及train/evaluate布尔开关、train_dataset/validation_dataset若为 None 会通过distribute_dataset(self.task.build_inputs, ...)从任务配置自动构建、checkpoint_exporter需提供maybe_export_checkpoint接口。另外训练/评估循环是否使用tf.function与tf.while_loop、是否启用 TPU summary 优化由config.trainer下的train_tf_function、train_tf_while_loop、eval_tf_function、eval_tf_while_loop、allow_tpu_summary字段控制见 base_trainer.py。三、第一步子类化 base Trainer文档给出的定制示例如下省略号与原文档一致表示省略部分与基类实现相同class CustomTrainer(base_trainer.Trainer): def __init__( self, config: ExperimentConfig, task: base_task.Task, model: tf.keras.Model, optimizer: tf.optimizers.Optimizer, train_dataset: Optional[Union[tf.data.Dataset, tf.distribute.DistributedDataset]] None, ...): super().__init__( configconfig, tasktask, modelmodel, optimizeroptimizer, train_datasettrain_dataset, ...) def train_step(self, iterator): def step_fn(inputs): if self.config.runtime.enable_xla and (self.config.runtime.num_gpus 0): task_train_step tf.function(self.task.train_step, jit_compileTrue) else: task_train_step self.task.train_step logs task_train_step(...) ... def eval_step(self, iterator): def step_fn(inputs): logs self.task.validation_step(...) ... return logs inputs, passthrough_logs self.next_eval_inputs(iterator) ... logs tf.nest.map_structure(...) return passthrough_logs | logs def train_loop_end(self): self.join() logs {} for metric in self.train_metrics [self.train_loss]: logs[metric.name] metric.result() metric.reset_states() if hasattr(self.optimizer, iterations): logs[learning_rate] self.optimizer.learning_rate( self.optimizer.iterations) ... logs[optimizer_iterations] self.optimizer.iterations logs[model_global_step] self.model._global_step return logs def eval_end(self, aggregated_logsNone): self.join() logs {} for metric in self.validation_metrics: logs[metric.name] metric.result() if self.validation_loss.count.numpy() ! 0: logs[self.validation_loss.name] self.validation_loss.result() ... if aggregated_logs: metrics self.task.reduce_aggregated_logs( aggregated_logs, global_stepself.global_step) logs.update(metrics) if self._checkpoint_exporter: self._checkpoint_exporter.maybe_export_checkpoint( self.checkpoint, logs, self.global_step.numpy()) ... return logs对照源码可以确认文档示例中的骨架与真实基类实现高度一致train_step的 XLA 分支源码中当config.runtime.enable_xla and config.runtime.num_gpus 0时用tf.function(self.task.train_step, jit_compileTrue)包裹训练步否则直接调用见 base_trainer.py随后self._train_loss.update_state(logs[self.task.loss])并self.global_step.assign_add(1)最后通过self.strategy.run(step_fn, args(inputs,), optionsself._runtime_options)在设备端执行。eval_step的 passthrough 日志机制next_eval_inputs返回的passthrough_logs与strategy.experimental_local_results聚合后的 logs 合并时模型日志键优先于 passthrough 键冲突时会打 warning见 base_trainer.py。eval_end中 EMAExponentialMovingAverage优化器的特殊处理eval_begin会swap_weights()切到平均权重做评估eval_end在“最优 checkpoint 导出之后”再swap_weights()切回以确保导出的是平均权重见 base_trainer.py 与 base_trainer.py。3.1 仓库中的真实示例RankingTrainer仓库内就有一个“最小化定制 Trainer”的完整范例official/recommendation/ranking/train.py 中的RankingTrainer。它只覆写了train_loop_end因为 Ranking 模型使用双优化器分别管理 embedding 与非 embedding 权重需要为每个优化器单独记录学习率class RankingTrainer(base_trainer.Trainer): A trainer for Ranking Model. The RankingModel has two optimizers for embedding and non embedding weights. Overriding train_loop_end method to log learning rates for each optimizer. def train_loop_end(self) - Dict[str, float]: self.join() logs {} for metric in self.train_metrics [self.train_loss]: logs[metric.name] metric.result() metric.reset_states() for i, optimizer in enumerate(self.optimizer.optimizers): lr_key f{type(optimizer).__name__}_{i}_learning_rate if callable(optimizer.learning_rate): logs[lr_key] optimizer.learning_rate(self.global_step) else: logs[lr_key] optimizer.learning_rate return logs这个例子展示了定制的最小闭环只继承、只覆写真正需要的方法其余流程checkpoint、summary、Orbit 循环完全复用基类。四、自定义训练驱动器launch script / Training driver4.1 动机official/vision/train.py 是 TFM 中启动模型训练的脚本。当你需要标准 Trainer 无法覆盖的功能时例如自定义训练循环 自定义启动流程的组合就需要基于它“分支”出定制的训练驱动器先创建自定义 Trainer再把它集成进自定义启动脚本通过run_experiment的trainer参数传入见 train_lib.py 的签名。文档还提到除 main 方法外可在驱动器中追加额外方法例如加载/保存模型权重、把训练进度写入日志文件、向指定通道发送训练进度通知等这些方法都可以从 main 中被调用。4.2 关键步骤一导入注册表Import the registry标准 train.py 中只有两行与注册表相关见 train.pyfrom official.vision import registry_imports # pylint: disableunused-import文档建议的写法是把所有自定义模型、任务、配置等“注册项”统一放进自己的registry_imports.py这样训练驱动器无需逐个文件处理from official import vision import registry_imports # pylint: disableunused-import其原理是TFM 的模型/任务/配置都通过registry.register注册到 official/core/registry.py而get_task、模型构建等工厂函数按名字查表。注册动作只发生在“模块被 import 时”所以驱动器必须在入口处 import 到这些模块否则会出现“注册表里找不到”的错误。仓库内的 official/vision/registry_imports.py 本身就是这样一个纯 import 聚合文件注释即写明 “All necessary imports for registration.”如果你的自定义模块较多建议照此模式单独建一个自定义 registry_imports 文件统一收纳。4.3 关键步骤二定义 main 方法文档对main方法的说明它是脚本入口负责编排整个训练流程main 内会调用 train_lib.run_experiment按实验参数执行 train 与 eval返回(model, eval_logs)二元组——eval_logs仅在run_post_eval为 True 时返回评估指标日志否则为{}这一点在 train_lib.py 的run()文档串中可以得到印证train_utils.save_gin_config 则负责序列化并保存本次实验的 gin 配置。文档给出的自定义启动脚本示例def main(_): ... if params.runtime.mixed_precision_dtype: performance.set_mixed_precision_policy(params.runtime.mixed_precision_dtype) distribution_strategy distribute_utils.get_distribution_strategy( dist_strategyparams.runtime.distribution_strategy, all_reduce_algparams.runtime.all_reduce_alg, num_gpusparams.runtime.num_gpus, tpu_addressparams.runtime.tpu) with distribution_strategy.scope(): task task_factory.get_task(params.task, logging_dirmodel_dir) ... train_lib.run_experiment( distribution_strategydistribution_strategy, tasktask, modeFLAGS.mode, paramsparams, model_dirmodel_dir) train_utils.save_gin_config(FLAGS.mode, model_dir) ... if __name__ __main__: tfm_flags.define_flags() flags.mark_flags_as_required([experiment, mode, model_dir]) app.run(main)4.4 与标准 train.py 的逐行对照把上面的示例与当前仓库实际的 official/vision/train.py 对照标准驱动器的完整 main 流程为gin.parse_config_files_and_bindings(FLAGS.gin_file, FLAGS.gin_params)先解析 gin 配置文件与命令行覆盖项params train_utils.parse_configuration(FLAGS)按--experiment指定的实验配置名构建ExperimentConfig若mode含train先train_utils.serialize_config(params, model_dir)把配置落盘为 yaml纯 eval 模式不写 yaml避免与持续评估任务写同一文件冲突见 train.py若配置了params.runtime.mixed_precision_dtypemixed_float16/mixed_bfloat16调用performance.set_mixed_precision_policy(...)设置混合精度进入_run_experiment_with_preemption_recoverytrain.py内部先distribute_utils.get_distribution_strategy(...)构造策略在strategy.scope()内task_factory.get_task(params.task, logging_dirmodel_dir)创建任务再调用train_lib.run_experiment(...)并传入eval_summary_manager与enable_async_checkpointing对应--enable_async_checkpointing旗标默认 True遇到 TPU 抢占tf.errors.OpError且带抢占信息时自动从最近 checkpoint 重启否则抛出异常最后train_utils.save_gin_config(FLAGS.mode, model_dir)保存 gin 配置。旗标约束与文档示例一致tfm_flags.define_flags()定义公共旗标flags.mark_flags_as_required([experiment, mode, model_dir])指定三个必填项见 train.py。从源码结构看mode支持的值由OrbitExperimentRunner.run决定train、train_and_post_eval、train_and_eval、eval、continuous_eval见 train_lib.py其中continuous_eval会以params.trainer.continuous_eval_timeout为超时、在global_step达到train_steps时退出。4.5 一个“自定义驱动器 自定义 Trainer”的完整参照如果你希望自定义 Trainer 真正生效驱动器必须显式构造它并通过trainer参数传给run_experiment。official/recommendation/ranking/train.py 给出了完整范式with strategy.scope(): checkpoint_exporter train_utils.maybe_create_best_ckpt_exporter( params, model_dir) trainer RankingTrainer( configparams, tasktask, modelmodel, optimizermodel.optimizer, traintrain in mode, evaluateeval in mode, train_datasettrain_dataset, validation_datasetvalidation_dataset, checkpoint_exportercheckpoint_exporter) train_lib.run_experiment( distribution_strategystrategy, tasktask, modemode, paramsparams, model_dirmodel_dir, trainertrainer)注意其中的两条关键细节Trainer 必须在strategy.scope()内创建——run_experiment的文档串也明确要求 “It should be created within the strategy.scope()”train_lib.py因为Trainer.__init__会通过tf.distribute.get_strategy()获取当前策略并据此分发数据集base_trainer.py数据集也可在驱动器内预先用strategy.distribute_datasets_from_function(get_dataset_fn(...), optionstf.distribute.InputOptions(experimental_fetch_to_deviceFalse))构造后传入从而跳过 Trainer 内部的自动构建路径。五、实践要点与适用前提选择定制层次只改训练循环细节 → 继承Trainer覆写单个钩子成本最低参考RankingTrainer需要改变启动流程额外旗标、额外回调、通知逻辑→ 复制 train.py 改 main两者组合时把自定义 Trainer 通过run_experiment(trainer...)注入即可。注册表先行自定义模型、任务、实验配置必须经由自定义registry_imports被 import否则task_factory.get_task与模型工厂查表失败。运行环境train_step中的jit_compileTrue分支仅在config.runtime.enable_xla且num_gpus 0时触发TPU 用户可另行通过config.runtime.tpu_enable_xla_dynamic_padder影响RunOptions见 base_trainer.py。配置可复现性main 中的serialize_config与save_gin_config两步会把实验参数与 gin 配置持久化到model_dir定制驱动器建议保留这两步便于实验回溯。相关文档延伸阅读同目录下的 customize_model_and_config.md、customize_input_pipeline.md、customize_training_process.md 分别覆盖模型/配置定制、输入管道定制与训练过程定制可与本文的“启动器定制”互为补充。综上TFM vision 的“启动器定制”本质上是两层可插拔设计的外化上层驱动器只负责“装配”gin 配置、分布式策略、任务工厂下层 Trainer 只负责“执行”Orbit 循环 钩子。沿着 train.py → train_lib.run_experiment → base_trainer.Trainer 这条链路定位差异点再选择“覆写钩子”或“分叉驱动器”的恰当组合即可在不破坏仓库既有训练框架的前提下完成绝大多数训练流程定制需求。【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/models创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价