首页
/ TensorFlow Model Garden 中的 DETR 端到端目标检测:从实验配置到匈牙利匹配的完整实现解析

TensorFlow Model Garden 中的 DETR 端到端目标检测:从实验配置到匈牙利匹配的完整实现解析

2026-09-04 13:27:23作者:傅爽业Veleda

本文围绕 TensorFlow Model Garden(tensorflow/models)中 DETR(End-to-End Object Detection with Transformers)项目的 README 及其配套源码展开:完整继承官方实验结果表与复现脚本,深入讲解 训练入口实验配置模型结构匈牙利匹配自定义 AdamW 优化器 的实现细节。读完本文,你可以按官方脚本复现 DETR-ResNet-50 在 COCO 上的训练流程,并理解从数据增强、匈牙利二部图匹配到损失归一化与模型导出的完整链路。

一、项目定位与实验结果总览

DETR 项目是论文《End-to-End Object Detection with Transformers》的 TensorFlow 2 实现,位于 official/projects/detr/ 目录。它把目标检测建模为一个集合预测问题:用一组可学习的 query 向量经过 Transformer 编解码器后直接输出固定数量的检测框,无需传统的 Anchor 生成、NMS 等非可微模块。

README 中给出的核心实验结果表如下(数据集为 COCO,权重图为 ImageNet 预训练的 ResNet-50,输入分辨率 1333x1333,全局 batch size 为 64):

模型 分辨率 Batch size Epochs Decay@ Params (M) Box AP 说明
DETR-ResNet-50 1333x1333 64 300 200 41 40.6 本仓库 300 轮复现(含 checkpoint 与实验记录)
DETR-ResNet-50 1333x1333 64 500 400 41 42.0 本仓库 500 轮复现(含 checkpoint 与实验记录)
DETR-ResNet-50 1333x1333 64 300 200 41 40.6 论文结果
DETR-ResNet-50 1333x1333 64 500 400 41 42.0 论文结果
DETR-DC5-ResNet-50 1333x1333 64 500 400 41 43.3 论文结果

从上表可以看出,本仓库 300/500 轮复现的 Box AP(40.6 / 42.0)与论文数值一致。每行实验对应 detr/experiments 下的一个脚本,预训练 checkpoint 托管在模型花园的 gs://tf_model_garden/vision/detr/ 路径下(如 detr_resnet_50_300.tar.gzdetr_resnet_50_500.tar.gz),训练曲线记录在 tensorboard.dev 的公开实验中。

README 同时在“Need contribution”一节明确列出了待办:DC5 支持尚未实现——即上表中 43.3 AP 的 DETR-DC5 版本目前只有论文结果,本仓库尚未提供对应代码,这是使用该实现时需要知晓的边界。

另外 README 保留了一条数据集免责声明:COCO、ImageNet 等数据集均由第三方提供并托管,使用相关数据前需阅读对应第三方平台的条款。

二、复现脚本与训练入口

两个实验脚本完整继承了 README 表格中的行设置:

300 轮脚本

#!/bin/bash
python3 official/projects/detr/train.py \
  --experiment=detr_coco \
  --mode=train_and_eval \
  --model_dir=/tmp/logging_dir/ \
  --params_override=task.init_checkpoint='gs://tf_model_garden/vision/resnet50_imagenet/ckpt-62400',trainer.train_steps=554400,trainer.optimizer_config.learning_rate.stepwise.boundaries="[369600]"

500 轮脚本 与之几乎相同,仅覆盖为:

python3 official/projects/detr/train.py \
  --experiment=detr_coco \
  --mode=train_and_eval \
  --model_dir=/tmp/logging_dir/ \
  --params_override=task.init_checkpoint='gs://tf_model_garden/vision/resnet50_imagenet/ckpt-62400'

两个脚本的共同点与差异点值得逐条解读:

  • --experiment=detr_coco:选择 configs/detr.py 中通过 @exp_factory.register_config_factory('detr_coco') 注册的实验工厂。该工厂默认按 500 轮 配置训练(train_steps = 500 * num_steps_per_epoch,衰减点在第 400 轮),因此 500 轮脚本无需额外覆盖;300 轮脚本则用 params_overridetrainer.train_steps 覆盖为 554400(即 300 轮),并把学习率衰减边界 boundaries 覆盖为 [369600](即 200 轮)。
  • task.init_checkpoint:指定 ImageNet 预训练的 ResNet-50 checkpoint。结合工厂默认 init_checkpoint_modules='backbone',加载时只恢复 backbone 权重(源码中使用 status.expect_partial() 做部分匹配),Transformer 主干从头训练。
  • --mode=train_and_eval--model_dir:分别指定训练+评估模式和输出目录。

train.py 作为统一训练驱动器,其主流程是 Model Garden 官方框架的标准四步(train.py#L35-L65):

  1. train_utils.parse_configuration(FLAGS) 解析出实验配置;训练模式下会把最终配置序列化为 YAML 写回 model_dir,便于复现与审计;
  2. 若配置了 runtime.mixed_precision_dtype,调用 performance.set_mixed_precision_policy 设置混合精度策略(GPU 场景下可用 float16 提速,TPU 场景下对应 bfloat16);
  3. 通过 distribute_utils.get_distribution_strategy 构建分布策略(单机多卡、TPU 等),并在其 scope() 内用 task_factory.get_task 依据 DetrTask 配置实例化任务类;
  4. train_lib.run_experiment 驱动训练/评估循环。

其中第 3 步依赖两个 # pylint: disable=unused-import 的“注册导入”:configs.detrtasks.detection。也就是说 detr_coco 实验名和 DetrTaskDetectionTask 的映射都是在模块导入时通过装饰器注册的,这是 Model Garden “配置工厂 + 任务工厂”解耦设计的关键机制。

三、实验配置:detr_coco 工厂逐项解读

configs/detr.py 定义了三组实验工厂和四类配置 dataclass。detr_coco(TFDS 版)工厂中的关键配置如下:

模型配置 DetrDetr dataclass,L55-L67):

参数 默认值 说明
num_queries 100 可学习 query 数量,即每图预测的固定检测框数上限
hidden_size 256 Transformer 隐层维度(位置编码特征数需与其相等且为偶数)
num_classes 91(工厂中覆盖为 81) 类别数含 background;COCO 实验中设为 81(80 类 + 背景)
num_encoder_layers / num_decoder_layers 6 / 6 与论文一致
input_size [1333, 1333, 3] 输入分辨率,短边多尺度 resize 后 pad 到 1333
backbone resnetmodel_id=50bn_trainable=False ResNet-50 且冻结 BN 统计量
backbone_endpoint_name '5' 取 ResNet 第 5 个 stage 的特征图作为 Transformer 输入

数据配置:训练集 tfds_name='coco/2017'tfds_split='train'shuffle_buffer_size=1000global_batch_size=64;验证集取 validation split 且 drop_remainder=False。常量 COCO_TRAIN_EXAMPLES = 118287COCO_VAL_EXAMPLES = 5000 被用于换算 num_steps_per_epoch 和训练总步数。

优化器配置detr_coco 工厂 L128-L145,三个工厂完全一致):

  • 优化器类型 detr_adamwweight_decay_rate=1e-4global_clipnorm=0.1,并显式设 gradient_clip_norm=0.0 以规避 AdamW 的 legacy 行为;
  • 学习率:stepwise 调度,boundaries=[decay_at]values=[1e-4, 1e-5],即在 decay_at 步从 1e-4 衰减到 1e-5,与论文一致;
  • 训练器层面:max_to_keep=1(仅保留最近一个 checkpoint),并以 AP 作为最佳 checkpoint 导出指标(best_checkpoint_eval_metric='AP',导出到 best_ckpt 子目录)。

此外还有两个变体工厂:detr_coco_tfrecord 面向本地 TFRecord 格式 COCO 数据(input_path='coco/train*',配合 annotation_file='coco/instances_val2017.json' 做评估),detr_coco_tfdsdetr_coco 类似但采用 Losses(class_offset=1) 的类别偏移处理。

四、模型结构:backbone + query + Transformer + 双头

核心模型类 DETR 定义在 modeling/detr.py。文件头注释明确指出:该模块不支持 Keras 序列化/反序列化,对象式保存应使用 tf.train.Checkpoint,图式导出应使用 tf.saved_model.save——这一点在导出章节会再次体现。

DETR 类(detr.py#L127-L262)的构建与推理流程如下:

  1. build 阶段build / _build_detection_decoder):

    • 一个 1x1 Conv2D_input_proj)把 backbone 第 5 个 stage 的特征投影到 hidden_size=256
    • query_embeddings:形状 [num_queries, hidden_size]、标准差为 1 的高斯初始化权重,即论文中的可学习 object queries;
    • 分类头 _class_embed:单个 Dense(num_classes)
    • 框回归头 _bbox_embed:两个 Dense(hidden_size, relu) 加一个 Dense(4),最后经 sigmoid 压回 [0, 1],输出相对坐标。
  2. call 前向detr.py#L225-L262):

    • 取 backbone 特征,由原始输入生成 padding mask(_generate_image_mask:对通道求和后判断非零,再最近邻缩放到特征图分辨率);
    • position_embedding_sine 生成 2D 正弦位置编码:对行、列坐标分别做 cumsum(从而自动跳过 padding 像素)、归一化到 [0, 2π] 后按 temperature=10000 展开为 sin/cos 特征并拼接(detr.py#L32-L95)。由于 mask 参与了坐标计算,位置编码对 padding 区域天然“感知不到”,这是该实现的细节亮点;
    • 特征图拉平为序列,与按 batch 复制的 query 一起送入 DETRTransformer
    • 解码器 return_all_decoder_outputs=True每一层解码器输出都会独立过一次分类头与框头。训练时对这些辅助输出逐一计算损失并累加(见下一节),推理时只取最后一个输出并调用 postprocess 转成标准检测输出。

postprocessdetr.py#L98-L124)把模型原始输出转成检测协议张量:

  • detection_boxesbox_ops.cycxhw_to_yxyx 把中心宽高的相对坐标转换为 yxyx 相对坐标;
  • detection_scores / detection_classes:对去掉 background 列[:, :, 1:])的类别 logits 取 softmax 最大值与 argmax(+1 对齐 COCO 1 起始的类别编号);
  • num_detections:源码注释标注该字段“暂未真正生效”,保留用于兼容。

Transformer 主干 DETRTransformerdetr.py#L265-L345)封装了 modeling/transformer.py 中的 TransformerEncoder / TransformerDecoder:8 头注意力、FFN 中间维度 2048、norm_first=False(post-norm,与论文一致)。调用时 encoder 以图像 mask 广播出的自注意力 mask 处理特征序列;decoder 的 self-attention mask 目前传入全 1 矩阵(源码注释引用了 b/199545430,说明是等待上游 bug 修复的临时写法),cross-attention 则以 query 数对图像 token 数广播同一 mask,query 与 memory 的位置编码分别传入 input_pos_embedmemory_pos_embed

五、损失函数与匈牙利匹配

训练任务类 DetectionTasktasks/detection.py)实现了 DETR 最核心的两部分:匈牙利匹配与三项损失。

匹配代价计算_compute_costdetection.py#L126-L170):

  • 分类代价:lambda_cls * -softmax(logits)[target],用负概率近似 1 - prob(省略不影响匹配的常数项);
  • 框 L1 代价:lambda_box * |pred - gt| 在 4 个坐标维上求和;
  • GIoU 代价:lambda_giou * (-GIoU),先把 cycxhw 相对框转 yxyx 再计算广义交并比;
  • 两个防御性处理:背景目标(cls_targets == 0)与 NaN/Inf 位置的代价一律置为 max_cost = lambda_cls + 4*lambda_box + lambda_giou 这一理论上限,避免 pad 位置参与匹配。

匈牙利匹配ops/matchers.pyhungarian_matching 纯 TensorFlow 实现。文件头注释说明其基于 Hungarian Matching Algorithm,思路是求解二部图最小权匹配。实现分为四段式原语:

  1. _prepare:按行、列最小值平移代价矩阵,使所有权重非负,得到“零元素”邻接矩阵作为贪心起点;
  2. _greedy_assignmenttf.foldl 逐行贪心分配(每行至多选一个未占用列);
  3. _find_augmenting_path + _improve_assignment:在贪心解基础上做增广路径搜索与回退翻转,tf.while_loop 迭代直到匹配最大;
  4. _compute_cover + _update_weights_using_cover:按 Kőnig 定理构造顶点覆盖,对未覆盖边减权、双重覆盖边加权,tf.while_loop 迭代更新权重直到覆盖完备。

整套算法全程以 3D 张量 [batch_size, num_elems, num_elems] 批处理,且各循环均 back_prop=False——匹配结果作为离散索引不反向传播,梯度只经由 stop_gradient 后的索引 gather 到对应 query 上。

损失组装build_lossesdetection.py#L172-L234):

  • 按匹配索引 gather 出被分配的分类 logits 与框输出;
  • 分类损失为带背景降权的交叉熵:背景项乘以 background_cls_weight=0.1 以缓解类别不平衡;
  • 框 L1 损失与 GIoU 损失只作用于非背景分配;
  • 归一化是分布式感知的:每个 replica 先统计本地框数与权重和,再用 replica_context.all_reduce(SUM) 汇聚到全局,损失除以全局框数/全局权重和。注释明确解释了这一设计——梯度聚合发生在优化器侧(replica sum),因此按全局框数归一化后损失值才可比、可解释;
  • 辅助损失 aux_losses(模型各解码层的额外损失,经 model.losses 注入)通过 tf.add_n 并入总损失。

train_stepdetection.py#L251-L321)对 model(features, training=True) 返回的每一层解码器输出循环调用 build_losses 并累加,支持 LossScaleOptimizer 的混合精度缩放路径,并在日志层面乘以 num_replicas_in_sync 还原可比较的数值。验证阶段(validation_step)只取最后一层输出,并把 predictions/ground_truths(含 source_idimage_infois_crowd 等)打包进 logs,交由 COCOEvaluator 计算 COCO AP。

六、数据管道:多尺度增强与坐标协议

DETR 的数据侧由两个 dataclass 配置驱动:

  • COCODataConfigdataloaders/coco.py):output_size=(1333, 1333)max_num_boxes=100resize_scales=(480, 512, ..., 768, 800) 共 11 档短边尺度;
  • 通用 DataConfig(configs/detr.py#L30-L42):dtype 默认 bfloat16shuffle_buffer_size=10000file_type='tfrecord'drop_remainder=True 等。

预处理逻辑(COCODataLoader.preprocesscoco.py#L42-L135;TFRecord 路径另有等价的 dataloaders/detr_input.py Parser)严格复刻了 DETR 论文的增强策略:

  1. 像素级 normalize_image(均值/标准差归一化),类别 label + 1 使 0 保留给 background;
  2. 训练时:随机水平翻转;以 50% 概率做随机裁剪增强——先把短边缩放到 {400, 500, 600} 之一,再随机切出边长在 [384, min(side, 600)] 内的 crop 并同步修正框坐标;
  3. 按训练 11 档 / 验证固定 800 的短边尺度 resize_image(长边不超过 1333),框经 resize_and_crop_boxesnormalize_boxes 转为相对坐标;
  4. 过滤全零框、yxyx_to_cycxhw 转换、pad 到 [1333, 1333],标签用 clip_or_pad_to_fixed_size 截断/补齐到 max_num_boxes=100——与 num_queries=100 一一对应;
  5. 验证时额外输出 idimage_info(缩放与 pad 信息,评估时用于还原绝对坐标)、is_crowdgt_boxes

批处理阶段(_transform_and_batch_fn)依据 input_context.get_per_replica_batch_size 把全局 batch 64 拆到各 replica,num_parallel_calls=tf.data.experimental.AUTOTUNE 并行化预处理。

七、优化器:给 backbone 打 0.1 倍学习率的定制 AdamW

optimization.py 注册了名为 detr_adamw 的定制优化器 _DETRAdamW(继承 official.nlp.optimization.AdamWeightDecay)。其与普通 AdamW 的唯一区别在 _resource_apply_dense / _resource_apply_sparseoptimization.py#L59-L149):

if 'detr' not in var.name:
  lr_t *= 0.1
...
lr = coefficients['lr_t'] * 0.1 if 'detr' not in var.name else coefficients['lr_t']

按变量名是否包含 detr 来区分模块:backbone(ImageNet 预训练权重,变量名不含 detr)使用 0.1 倍学习率微调,Transformer 主干与两个检测头使用全量学习率。这是一种轻量级的分层学习率策略,配合 weight_decay_rate=1e-4global_clipnorm=0.1 的全局梯度裁剪,共同构成 README 中“Customized optimizer to match paper results”(文件头 docstring 原话:定制优化器以匹配论文结果)的说明。

八、预训练权重加载与模型导出

backbone 预热DetectionTask.initializedetection.py#L67-L88)实现 init_checkpoint 的两种粒度——'all' 时全量 assert_consumed'backbone'detr_coco 工厂的默认值)时只构造 tf.train.Checkpoint(backbone=model.backbone)expect_partial(),即允许 checkpoint 中多余的 BN 统计量等变量不被消费。脚本中的 gs://tf_model_garden/vision/resnet50_imagenet/ckpt-62400 即按此路径加载。

导出与服务serving/export_module.py 中的 DETRModule 继承通用 DetectionModule_build_model 按相同配置重建模型并用 tf_keras.Input 触发构建;serve 入口接受 uint8 [batch, None, None, 3] 图像,在 CPU 上完成均值/方差归一化与短边 1333 缩放,取最后一层解码器输出并还原 detection_boxes 为绝对像素坐标;若 input_type='tflite' 则跳过图内预处理以兼容 TFLite 量化。配套的 export_saved_model.py 提供 SavedModel 导出入口。如 modeling/detr.py 文件头所述,该模型不支持 Keras 反序列化,对象式存取请用 tf.train.Checkpoint,图式序列化请用 tf.saved_model.save,导出与加载环节需遵守这一约定。

九、运行前提与引用

运行前提:代码依赖 Model Garden 官方框架(official.core 配置/训练体系)与 TensorFlow 2 生态,依赖清单见 requirementsnightly_requirements;训练需要 COCO 2017 数据集(TFDS 或 TFRecord 格式)与若干百 GB 量级的磁盘空间(1333x1333 输入、64 全局 batch、300~500 轮),README 提供的 checkpoint 可直接用于对照或继续微调。

若该代码库对你的研究有帮助,README 建议引用 TensorFlow Model Garden:

@misc{tensorflowmodelgarden2020,
  author = {Hongkun Yu and Chen Chen and Xianzhi Du and Yeqing Li and
            Abdullah Rashwan and Le Hou and Pengchong Jin and Fan Yang and
            Frederick Liu and Jaeyoun Kim and Jing Li},
  title = {{TensorFlow Model Garden}},
  howpublished = {\url{https://github.com/tensorflow/models}},
  year = {2020}
}

十、小结

  • 实验可复现性detr/experiments 下两个脚本与 40.6/42.0 AP 的复现结果一一对应,params_override 精确覆盖了轮数与衰减点,detr_coco 工厂默认即 500 轮论文配置;
  • 架构关键点:ResNet-50 第 5 stage + 1x1 投影、100 个可学习 query、6+6 层 post-norm Transformer、每层解码器输出都参与辅助损失;
  • 工程亮点:纯 TF 实现的匈牙利匹配(批处理、可微索引 + stop_gradient)、分布式感知的全局框数归一化、按变量名分层学习率的定制 AdamW;
  • 已知边界:DC5 变体尚未实现(README 明确的贡献缺口),num_detections 字段暂未真正生效,decoder self-attention mask 存在临时写法。
登录后查看全文
热门项目推荐
相关项目推荐