首页
/ TensorFlow Model Garden 技术指南:Official 与 Research 模型、训练实验框架与 Orbit 训练循环

TensorFlow Model Garden 技术指南:Official 与 Research 模型、训练实验框架与 Orbit 训练循环

2026-09-03 16:08:05作者:谭伦延

本篇基于 TensorFlow Models 仓库的总览文档(docs/index.md)展开,系统讲解 Model Garden 的两大模型类别(Official / Research)、声明式训练实验框架的三级配置对象(runtime / task / trainer)及其真实 YAML 配置示例,以及 Orbit 训练循环管理库的定位与源码实现。读完后,你将能够理解 Model Garden 的整体资源版图,掌握用 YAML 或 Python 声明式配置快速拉起一次 ImageNet 训练的方法,并弄清 Orbit 在 Keras Model.fit 与手写训练循环之间的定位及其内部机制。

Model Garden 是什么

TensorFlow Model Garden 提供大量计算机视觉(Vision)与自然语言处理(NLP)领域最先进(state-of-the-art)机器学习模型的实现,以及配套的工作流工具,让你可以快速配置并在标准数据集上运行这些模型。无论是想为某个知名模型做性能基准测试(benchmark)、复现近期发表的研究结果,还是在现有模型基础上做扩展,Model Garden 都能支撑 ML 研究与工程应用。

仓库提供的资源面向 ML 开发者,包含五类:

  • Official models(官方模型):面向视觉与 NLP,由 Google 工程师维护;
  • Research models(研究模型):随 ML 研究论文一同公开的模型实现;
  • Training experiment framework(训练实验框架):用声明式方式快速配置并运行官方模型的训练实验;
  • Specialized ML operations(专用 ML 算子库):面向视觉与 NLP 的专用算子;
  • Model training loop 管理库 Orbit:用于管理自定义训练循环。

这些资源构建在 TensorFlow Core 框架之上,可与既有的 TensorFlow 开发项目集成,并以 Apache 开源协议发布,可以自由扩展与分发。

一个重要的工程事实是:实用的 ML 模型训练与推理在计算上都很重,通常需要 GPU、TPU 等加速器。Model Garden 中大多数模型是在 TPU 上基于大规模数据集训练的,但同样支持在 GPU 与 CPU 上训练和运行——这一点在后续的配置示例中会直接体现(同一模型家族通常同时提供 _tpu.yaml_gpu.yaml 两套配置)。

两大模型类别:Official 与 Research

Model Garden 中的模型都附带完整代码,可供测试、训练与再训练(re-train),用于研究与实验。仓库将其分为两大类,二者在维护方、API 版本和支持方式上有明确区别。

Official Models:面向 TF 2.x 的高层 API 实现

Official Models 目录是一个聚焦视觉与 NLP 的最先进模型集合,其实现基于当前 TensorFlow 2.x 高层 API。这些模型库被优化以获得较快的性能,并由 Google 工程师积极维护。官方模型还附带额外元数据,可以直接用于 Model Garden 的训练实验框架来快速配置实验。

official/README.md 的模型清单可以看到其覆盖范围:

计算机视觉

模型 类别
ResNet / ResNet-RS 图像分类
EfficientNet 图像分类
Vision Transformer (ViT) 图像分类
RetinaNet 目标检测
Mask R-CNN 实例分割
YOLO 目标检测(位于 official/projects/yolo
SpineNet 目标检测
MoViNets 视频分类(位于 official/projects/movinet

自然语言处理

模型 类别
ALBERT / BERT 预训练语言模型
ELECTRA 预训练语言模型
Transformer 神经机器翻译
NHNet 新闻标题生成
MobileBERT 知识蒸馏

推荐系统

模型 说明
DLRM 深度推荐模型
DCN v2 Web 级学习排序
NCF 神经协同过滤

官方 README 还说明了使用前提:master 分支的官方模型基于 TensorFlow 2 master 分支开发,克隆仓库或安装 pip 二进制包时会拉取 master 分支的 TensorFlow 作为依赖,等价于:

pip3 install tf-models-nightly
pip3 install tensorflow-text-nightly # 当模型使用 `nlp` 包时需要

Research Models:随论文公开的代码资源

Research Models 目录收录的是作为论文代码资源公开发布的模型实现,同时使用 TensorFlow 1.x 与 2.x 实现。与研究模型库的支持方式不同,这些代码由各自的代码作者与研究社区维护(maintained by their respective authors),而非 Google 官方团队。

research/README.md 可以看到其组织方式——每个模型目录都标注了对应论文、发表会议与维护者:

  • 建模库object_detection(TensorFlow Object Detection API,附带 COCO、KITTI、Open Images 等预训练模型)、slim(图像分类模型库,含 Inception、ResNet、VGG、MobileNet、NASNet 等);
  • 计算机视觉attention_ocrautoaugment(AutoAugment / Shake-Shake / ShakeDrop)、deeplab(DeepLab v1~v3+)、delf(大规模图像检索)、lstm_object_detectionvid2depth
  • NLPadversarial_text(半监督文本对抗训练)、cvt_text(Cross-View Training);
  • 音频与语音audioset(VGGish / YAMNet)、deep_speech(Deep Speech 2);
  • 强化学习efficient-hrlpcl_rl
  • 其他lfadsrebar 等。

因此选择模型时的判断依据很清晰:需要长期维护、高性能和声明式训练配置,用 Official;需要复现某篇具体论文的实现细节,去 Research 找对应目录

训练实验框架:声明式配置快速拉起训练

Model Garden 的训练实验框架(training experiment framework)让你可以借助官方模型与标准数据集快速组装并运行训练实验。该框架利用了官方模型自带的额外元数据,允许以声明式编程模型快速配置模型:既可以用 Python 命令在 TensorFlow Model 库(tfm 包,即 official 目录提供的 tfm.core 等模块)中定义实验,也可以用 YAML 配置文件。

三级配置对象:runtime / task / trainer

训练框架以 tfm.core.base_trainer.ExperimentConfig 作为顶层配置对象。在仓库源码中它定义于 official/core/config_definitions.py,是一个只含三个字段的数据类:

@dataclasses.dataclass
class ExperimentConfig(base_config.Config):
  """Top-level configuration."""
  task: TaskConfig = dataclasses.field(default_factory=TaskConfig)
  trainer: TrainerConfig = dataclasses.field(default_factory=TrainerConfig)
  runtime: RuntimeConfig = dataclasses.field(default_factory=RuntimeConfig)

三个顶层配置对象的职责划分是:

  • runtime(RuntimeConfig):定义处理硬件、分布式策略与其他性能优化。从源码看,其关键字段包括 distribution_strategy(如 'mirrored''tpu',默认 'mirrored')、enable_xlatpu(TPU 地址)、num_gpus(GPU 数量)、worker_hosts / task_index(多 worker 训练)、mixed_precision_dtype'float32' / 'float16' / 'bfloat16')、loss_scale,以及模型并行相关的 num_cores_per_replicadefault_shard_dimuse_tpu_mp_strategy 等;
  • task(TaskConfig):定义模型、训练数据、损失函数与初始化。它包含 init_checkpointmodeltrain_datavalidation_data(后两者均为 DataConfig)等字段,另有面向差异化隐私(differential privacy)与图像 summary 的可选配置;
  • trainer(TrainerConfig):定义优化器、训练循环、评估循环、summaries 与 checkpoints。关键默认值包括:steps_per_loop = 1000summary_interval = 1000checkpoint_interval = 1000max_to_keep = 5train_steps = 0validation_steps = -1(表示评估整个数据集)、validation_interval = 1000,以及优化器配置 optimizer_config(内含优化器、学习率与 warmup 调度)。

其中 DataConfig(同样定义在 official/core/config_definitions.py)值得注意的字段有:input_path(TFRecord 文件路径或逗号分隔的多个路径)与 tfds_name / tfds_split(直接从 TensorFlow Datasets 加载,二者互斥)、global_batch_size(跨所有副本的全局 batch size)、is_trainingdrop_remainder(默认 True)、shuffle_buffer_sizecacheenable_tf_data_service 等。

真实 YAML 示例:ImageNet ResNet50 on TPU

原文档给出的示例是 ResNet50 的 TPU 配置,其完整内容见 imagenet_resnet50_tpu.yaml

runtime:
  distribution_strategy: 'tpu'
  mixed_precision_dtype: 'bfloat16'
task:
  model:
    num_classes: 1001
    input_size: [224, 224, 3]
    backbone:
      type: 'resnet'
      resnet:
        model_id: 50
  losses:
    l2_weight_decay: 0.0001
    one_hot: true
    label_smoothing: 0.1
  train_data:
    input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/train*'
    is_training: true
    global_batch_size: 4096
    dtype: 'bfloat16'
  validation_data:
    input_path: 'gs://mlcompass-data/imagenet/imagenet-2012-tfrecord/valid*'
    is_training: false
    global_batch_size: 4096
    dtype: 'bfloat16'
    drop_remainder: false
trainer:
  train_steps: 28080
  validation_steps: 13
  validation_interval: 312
  steps_per_loop: 312
  summary_interval: 312
  checkpoint_interval: 312
  optimizer_config:
    optimizer:
      type: 'sgd'
      sgd:
        momentum: 0.9
    learning_rate:
      type: 'stepwise'
      stepwise:
        boundaries: [9360, 18720, 24960]
        values: [1.6, 0.16, 0.016, 0.0016]
    warmup:
      type: 'linear'
      linear:
        warmup_steps: 1560

对照 config_definitions.py 逐项解读这份配置:

  • runtime 段只用了两个字段:distribution_strategy: 'tpu' 指定在 TPU 上以 TPU 分布式策略运行;mixed_precision_dtype: 'bfloat16' 指定混合精度。TPU 上习惯用 bfloat16(动态范围与 float32 相同,适合大 batch),这与训练数据中 dtype: 'bfloat16' 相互呼应;
  • task.model 通过 backbone.type: 'resnet' + model_id: 50 从注册表实例化 ResNet-50 骨干,num_classes: 1001 是 ImageNet-2012 的类别数,input_size: [224, 224, 3] 为输入张量形状;
  • task.lossesl2_weight_decay: 0.0001 是权重衰减,label_smoothing: 0.1 是标签平滑,one_hot: true 表示标签为 one-hot 编码;
  • trainer.optimizer_config 展示了 OptimizationConfig 的三段式结构:optimizer(SGD + 0.9 动量)、learning_rate(stepwise 分段学习率,在 9360 / 18720 / 24960 步处从 1.6 依次衰减到 0.0016)、warmup(前 1560 步线性 warmup)。train_steps: 28080 恰好是 12 * 9360 * ... 量级的 90 epoch 配置(4096 全局 batch 下约 1251 步/epoch),三个 boundary 正好落在 1/3、2/3、3/4 训练进度处——这是 ResNet 论文中典型的阶梯式衰减节奏;
  • validation_steps: 13drop_remainder: false 搭配使用:13 × 4096 ≈ 53248,即整个 ImageNet 验证集按 batch 4096 评估一遍;312 这一数值(steps_per_loop / summary_interval / checkpoint_interval / validation_interval)与 28080 的整除关系保证了每个 loop 结束时正好触发评估与检查点。

同一目录下还有 48 个同类实验配置(imagenet_mobilenetv1_tpu.yamlimagenet_vitb16_i224_gpu.yamlimagenet_resnetrs50_i160_gpu.yaml 等),覆盖 ResNet、ResNet-RS、MobileNet v3/v4、ViT、DeepLab 等,且同一模型普遍同时存在 _tpu_gpu 变体,直接印证了"TPU 训练、GPU/CPU 亦可运行"的说明。

从配置到训练:训练驱动的源码调用链

以视觉任务的训练入口 official/vision/train.py 为例,从源码结构看,YAML 配置最终是这样被消费的:入口脚本用 absl flags 接收 --config_path 等命令行参数后,经由 official/core 中的通用训练库执行实验,核心调用链为:

  1. distribute_utils.get_distribution_strategy(...):根据 params.runtime.distribution_strategynum_gpustpu 地址等字段构建 tf.distribute 策略;
  2. task_factory.get_task(params.task, ...):在 distribution_strategy.scope() 内,依据 task 配置实例化对应任务(模型 + 数据管线 + 损失 + 指标);
  3. train_lib.run_experiment(distribution_strategy=..., task=..., ...):执行训练/评估循环,循环内的 checkpoint、summary、评估等行为由 TrainerConfig 的间隔字段驱动。

入口脚本还实现了断点恢复逻辑(_run_experiment_with_preemption_recovery):在 TPU 等可能被抢占的环境中捕获中断并尝试重连续训,这与 TrainerConfigpreemption_on_demand_checkpoint(被抢占时保存按需检查点,默认 True)的配置项相配合。

对于想走 Python 而非 YAML 路线的用户,仓库提供完整教程 Image classification with Model Garden,演示如何用 tfm 库在代码中定义并运行实验。

专用 ML 算子库(Specialized ML operations)

Model Garden 包含大量专为视觉与 NLP 设计的算子,目标是让最先进模型在 GPU 与 TPU 上高效执行。对应的源码位于:

  • 视觉算子official/vision/ops/ 目录,包含约 20 个算子模块(如卷积、归一化、池化、空间变换等),配合 official/vision/modelingofficial/vision/dataloaders 完成模型搭建与数据管线;
  • NLP 算子official/nlp/modeling/ops/ 目录,服务于 Transformer 类模型的训练与执行。

这些库除了核心算子外,还包含视觉与 NLP 数据处理、训练和模型执行所需的辅助函数。相关 API 的完整清单以仓库源码为准(tfm.visiontfm.nlp 命名空间),本文不重复罗列。如果你希望直观看到这些算子与任务如何被组合成端到端流程,可参考 docs/vision 下的四个教程:image_classification.ipynbobject_detection.ipynbinstance_segmentation.ipynbsemantic_segmentation.ipynb

训练循环管理:Orbit 的定位与实现

训练 TensorFlow 模型时,默认有两条路:

  1. Keras 高层 Model.fit:如果你的模型和训练流程符合 Keras Model.fit 的假设(对数据 batch 做增量梯度下降),它非常方便;
  2. 手写训练循环:用 tf.GradientTapetf.function 等低层 API 从零实现。这种方式灵活,但样板代码很多,且不会帮你简化分布式训练。

Orbit 就是在这两个极端之间提供的第三种选择。它是一个灵活、轻量的库,专为在 TensorFlow 2.x 中编写自定义训练循环而设计,并与 Model Garden 训练实验框架配合良好。Orbit 负责处理常见的训练事务——保存 checkpoint、运行模型评估、设置 summary 写入——同时把内层训练循环(inner training loop)的实现权完全交给用户。它无缝集成 tf.distribute,支持在 CPU、GPU 与 TPU 上运行,核心代码以易读、易 fork 为目标,并以 开源协议 发布(见 orbit/README.md)。

源码视角:StandardTrainer 的循环结构

orbit/standard_runner.py 的模块 docstring 与类实现可以看到 Orbit 的分层设计:

  • orbit/runner.py 定义抽象接口 AbstractTrainer / AbstractEvaluator,是最底层的循环契约;
  • StandardTrainer / StandardEvaluator 在其上增加结构:把训练/评估循环拆成 train_loop_begin() → 循环体 train_step(train_iterator)train_loop_end() 三段,并由 Orbit 提供循环本身的实现;
  • StandardTrainer 支持用 tf.while_loop 结构运行循环以获得额外性能(在 TPU 上尤其明显),并提供 TPU summary 写入的性能优化。

控制这些行为的 StandardTrainerOptions(frozen dataclass)在源码中给出了三个开关及默认值:

选项 默认值 作用
use_tf_function True 对循环体(涉及 train_step)应用 tf.functiontrain_loop_begin / train_loop_end 始终以 eager 模式运行
use_tf_while_loop True tf.while_loop 运行训练循环(要求同时 use_tf_function=True
use_tpu_summary_optimization False TPU 上条件式写 summary 极慢时启用:创建两个 XLA 程序(一个含 summary 调用、一个不含),含 summary 的程序仅在需要记录的那一步运行

这解释了 TrainerConfigtrain_tf_while_looptrain_tf_functioneval_tf_functioneval_tf_while_loop 等字段的用途——它们正是把 Orbit 的 StandardTrainerOptions 映射到了声明式配置里;allow_tpu_summary 则对应 TPU summary 优化开关。当 StandardTrainer 无法满足需求时(例如 StandardEvaluator 不支持对多个不同评估数据集做完整评估),文档明确建议用户直接回退到自定义 AbstractTrainer / AbstractEvaluator 子类。

Orbit 的完整使用指南见 docs/orbit/index.ipynb,其中演示了如何基于 StandardTrainer 实现自己的任务与数据管线;orbit/examples/single_task 目录还附有单任务最小示例代码。

与 Keras 自定义训练的关系

原文档同时指出:Keras API 本身的训练行为也是可定制的——通常通过覆写 Model.train_step 方法,或使用 callbacks.ModelCheckpointcallbacks.TensorBoardkeras.callbacks。换言之,Orbit 面向的场景是"需要完全掌控内层循环(例如非标准优化过程、TPU while-loop 性能优化),又不想手写全部样板代码"的中间地带;如果你只是常规微调,Model.fit 加少量 callback 往往更省事。

小结与使用建议

  • 选模型:追求维护质量、性能与声明式配置 → official;复现特定论文 → research
  • 跑实验:优先找现成 YAML(official/*/configs/experiments/ 下有 80 余个视觉实验配置),按 imagenet_resnet50_tpu.yaml 的模式修改 runtime / task / trainer 三段即可;
  • 配置含义存疑:对照 official/core/config_definitions.py 中各 dataclass 的字段文档,那里写明了每个字段的取值、默认值与相互约束(如 input_pathtfds_name 互斥、validation_steps=-1 表示评估全量数据);
  • 训练循环不满足需求:在训练实验框架与 Orbit 之间二选一——前者管"配置驱动的一整套实验",后者管"循环内每一步的细节";
  • 硬件:绝大多数实验配置面向 TPU 编写(distribution_strategy: 'tpu'bfloat16),但仓库为同一模型普遍提供 GPU 变体,CPU 亦可运行小规模实验;安装时注意官方模型依赖 TensorFlow master(nightly)版本这一前提。
登录后查看全文
热门项目推荐
相关项目推荐