TensorFlow Model Garden 中的 DETR 端到端目标检测:从实验配置到匈牙利匹配的完整实现解析
本文围绕 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.gz、detr_resnet_50_500.tar.gz),训练曲线记录在 tensorboard.dev 的公开实验中。
README 同时在“Need contribution”一节明确列出了待办:DC5 支持尚未实现——即上表中 43.3 AP 的 DETR-DC5 版本目前只有论文结果,本仓库尚未提供对应代码,这是使用该实现时需要知晓的边界。
另外 README 保留了一条数据集免责声明:COCO、ImageNet 等数据集均由第三方提供并托管,使用相关数据前需阅读对应第三方平台的条款。
二、复现脚本与训练入口
两个实验脚本完整继承了 README 表格中的行设置:
#!/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_override把trainer.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):
train_utils.parse_configuration(FLAGS)解析出实验配置;训练模式下会把最终配置序列化为 YAML 写回model_dir,便于复现与审计;- 若配置了
runtime.mixed_precision_dtype,调用performance.set_mixed_precision_policy设置混合精度策略(GPU 场景下可用 float16 提速,TPU 场景下对应 bfloat16); - 通过
distribute_utils.get_distribution_strategy构建分布策略(单机多卡、TPU 等),并在其scope()内用task_factory.get_task依据DetrTask配置实例化任务类; train_lib.run_experiment驱动训练/评估循环。
其中第 3 步依赖两个 # pylint: disable=unused-import 的“注册导入”:configs.detr 与 tasks.detection。也就是说 detr_coco 实验名和 DetrTask 到 DetectionTask 的映射都是在模块导入时通过装饰器注册的,这是 Model Garden “配置工厂 + 任务工厂”解耦设计的关键机制。
三、实验配置:detr_coco 工厂逐项解读
configs/detr.py 定义了三组实验工厂和四类配置 dataclass。detr_coco(TFDS 版)工厂中的关键配置如下:
模型配置 Detr(Detr 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 |
resnet,model_id=50,bn_trainable=False |
ResNet-50 且冻结 BN 统计量 |
backbone_endpoint_name |
'5' |
取 ResNet 第 5 个 stage 的特征图作为 Transformer 输入 |
数据配置:训练集 tfds_name='coco/2017'、tfds_split='train'、shuffle_buffer_size=1000、global_batch_size=64;验证集取 validation split 且 drop_remainder=False。常量 COCO_TRAIN_EXAMPLES = 118287、COCO_VAL_EXAMPLES = 5000 被用于换算 num_steps_per_epoch 和训练总步数。
优化器配置(detr_coco 工厂 L128-L145,三个工厂完全一致):
- 优化器类型
detr_adamw:weight_decay_rate=1e-4、global_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_tfds 与 detr_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)的构建与推理流程如下:
-
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],输出相对坐标。
- 一个
-
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转成标准检测输出。
- 取 backbone 特征,由原始输入生成 padding mask(
postprocess(detr.py#L98-L124)把模型原始输出转成检测协议张量:
detection_boxes:box_ops.cycxhw_to_yxyx把中心宽高的相对坐标转换为 yxyx 相对坐标;detection_scores/detection_classes:对去掉 background 列([:, :, 1:])的类别 logits 取 softmax 最大值与 argmax(+1对齐 COCO 1 起始的类别编号);num_detections:源码注释标注该字段“暂未真正生效”,保留用于兼容。
Transformer 主干 DETRTransformer(detr.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_embed 与 memory_pos_embed。
五、损失函数与匈牙利匹配
训练任务类 DetectionTask(tasks/detection.py)实现了 DETR 最核心的两部分:匈牙利匹配与三项损失。
匹配代价计算(_compute_cost,detection.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.py 的 hungarian_matching 纯 TensorFlow 实现。文件头注释说明其基于 Hungarian Matching Algorithm,思路是求解二部图最小权匹配。实现分为四段式原语:
_prepare:按行、列最小值平移代价矩阵,使所有权重非负,得到“零元素”邻接矩阵作为贪心起点;_greedy_assignment:tf.foldl逐行贪心分配(每行至多选一个未占用列);_find_augmenting_path+_improve_assignment:在贪心解基础上做增广路径搜索与回退翻转,tf.while_loop迭代直到匹配最大;_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_losses,detection.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_step(detection.py#L251-L321)对 model(features, training=True) 返回的每一层解码器输出循环调用 build_losses 并累加,支持 LossScaleOptimizer 的混合精度缩放路径,并在日志层面乘以 num_replicas_in_sync 还原可比较的数值。验证阶段(validation_step)只取最后一层输出,并把 predictions/ground_truths(含 source_id、image_info、is_crowd 等)打包进 logs,交由 COCOEvaluator 计算 COCO AP。
六、数据管道:多尺度增强与坐标协议
DETR 的数据侧由两个 dataclass 配置驱动:
COCODataConfig(dataloaders/coco.py):output_size=(1333, 1333)、max_num_boxes=100、resize_scales=(480, 512, ..., 768, 800)共 11 档短边尺度;- 通用
DataConfig(configs/detr.py#L30-L42):dtype默认bfloat16、shuffle_buffer_size=10000、file_type='tfrecord'、drop_remainder=True等。
预处理逻辑(COCODataLoader.preprocess,coco.py#L42-L135;TFRecord 路径另有等价的 dataloaders/detr_input.py Parser)严格复刻了 DETR 论文的增强策略:
- 像素级
normalize_image(均值/标准差归一化),类别label + 1使 0 保留给 background; - 训练时:随机水平翻转;以 50% 概率做随机裁剪增强——先把短边缩放到 {400, 500, 600} 之一,再随机切出边长在
[384, min(side, 600)]内的 crop 并同步修正框坐标; - 按训练 11 档 / 验证固定 800 的短边尺度
resize_image(长边不超过 1333),框经resize_and_crop_boxes与normalize_boxes转为相对坐标; - 过滤全零框、
yxyx_to_cycxhw转换、pad 到[1333, 1333],标签用clip_or_pad_to_fixed_size截断/补齐到max_num_boxes=100——与num_queries=100一一对应; - 验证时额外输出
id、image_info(缩放与 pad 信息,评估时用于还原绝对坐标)、is_crowd、gt_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_sparse(optimization.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-4 与 global_clipnorm=0.1 的全局梯度裁剪,共同构成 README 中“Customized optimizer to match paper results”(文件头 docstring 原话:定制优化器以匹配论文结果)的说明。
八、预训练权重加载与模型导出
backbone 预热:DetectionTask.initialize(detection.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 生态,依赖清单见 requirements 与 nightly_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 存在临时写法。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0623
Hy4-previewHy4 preview 是由腾讯混元团队研发的新一代混合专家(MoE)旗舰模型。模型总参数量 770B,每个 token 激活 49B,主干共包含78层,第一层采用标准 FFN,其余 77 层均为 MoE 结构,每层包含 256 个路由专家与 1 个共享专家,每个 token 激活 top-8 路由专家及共享专家。主干之外原生内置 1 层 MTP(总参数量 10B,激活 0.7B)以支持投机解码。Python00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
GLM-5.3-FlashGLM-5.3-Flash (320B-A18B),是GLM-5系列的首个原生多模态模型。320B总参数,能力超过GLM-5.2Jinja00
Spark-X2.5-4BSpark-X2.5-4B 旨在让强大的 AI 更实用、更高效、更易获得。在广泛日常任务中表现强劲,涵盖对话、写作、翻译、推理、编码、工具调用以及智能体工作流,并在同等规模的开源模型中取得领先成绩。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00