MMDetection 多目标跟踪模型训练与测试完全指南:QDTrack / SORT / DeepSORT / Mask2Former VIS 实战
2026/9/19 14:09:19 网站建设 项目流程
  • 人工智能
  • 计算机视觉
  • 深度学习
  • 模型评测

【免费下载链接】mmdetection

OpenMMLab Detection Toolbox and Benchmark

项目地址:https://gitcode.com/gh_mirrors/mm/mmdetection
点击查看免费下载

本文档是 MMDetection 仓库中多目标跟踪(MOT)与视频实例分割(VIS)模型训练与测试的实战指南,覆盖 CPU、单 GPU、单节点多 GPU、多节点与 Slurm 五种运行环境下的完整操作流程。读完本文,你将掌握如何使用 tools/train.py 训练 QDTrack 等跟踪模型,如何使用 tools/test_tracking.py 配合--checkpoint/--detector/--reid参数完成 SORT、DeepSORT、StrongSORT、ByteTrack、OCSORT、Mask2Former VIS 等模型的测试与评估,并理解val_interval、学习率线性扩展规则、--amp混合精度、TrackImgSampler图像级采样、format_only结果格式化等关键机制背后的源码实现。

背景:为什么跟踪模型需要独立的训练与测试入口

MMDetection 中的跟踪任务(tracking)与传统检测任务(detection)在数据组织、采样方式和评估指标上有本质差异:跟踪数据以「视频/序列」为单位,评估指标采用 HOTA / CLEAR / Identity 等跟踪专用指标(详见 mot_challenge_metric.py),部分算法还需要同时加载检测器与 ReID 模型两套权重。因此,训练使用通用入口 tools/train.py,而测试则使用独立入口 tools/test_tracking.py,并配套 tools/dist_train.sh、tools/dist_test_tracking.sh、tools/slurm_train.sh、tools/slurm_test_tracking.sh 四个分布式/集群脚本。

一、训练现有跟踪模型

本节介绍如何在受支持的数据集上训练现有跟踪模型,支持以下训练环境:

  • CPU
  • 单 GPU
  • 单节点多 GPU
  • 多节点

此外,在 Slurm 管理的集群上,可以使用slurm_train.sh提交训练作业。

1.1 训练前须知(重要约定)

通过train_cfg调整评估间隔

训练过程中,可以通过修改配置中的train_cfg来控制评估(validation)频率。例如:

train_cfg = dict(val_interval=10)

表示每训练 10 个 epoch 对模型评估一次。该配置在 MMEngine Runner 的训练循环中被解析,是调节「训练耗时」与「指标监控粒度」之间平衡的关键参数。

学习率按线性扩展规则缩放

所有配置文件中的默认学习率均按 8 个 GPU 设定。根据线性扩展规则(Linear Scaling Rule,出自 "Accurate, Large Minibatch SGD: Training ImageNet in 1 Hour"),当每个 GPU 上的图像数或 GPU 总数变化时,学习率应与批次大小成比例调整:

  • 8 个 GPU × 每 GPU 1 张图像,学习率为lr=0.01
  • 16 个 GPU × 每 GPU 2 张图像,学习率为lr=0.04(批次扩大 4 倍,学习率同步扩大 4 倍)。

除手动换算外,tools/train.py 还提供了--auto-scale-lr参数:当配置中存在auto_scale_lr.enableauto_scale_lr.base_batch_size字段时,它会自动按当前实际批次大小缩放学习率,否则会抛出RuntimeError提示配置缺失。

日志与检查点保存位置

训练过程中的日志文件和检查点(checkpoint)将保存到工作目录,该目录由 CLI 参数--work-dir指定。默认值为./work_dirs/CONFIG_NAME(CONFIG_NAME 即去掉.py后缀的配置文件名)。

从 tools/train.py 的源码可以看到 work_dir 的确定优先级为:CLI 参数 > 配置文件中的work_dir字段 > 配置文件文件名。也就是说,命令行传入的--work-dir优先级最高,其次才是配置内定义的工作目录。

混合精度训练

如果需要混合精度训练,只需在命令行指定--amp参数。源码中该参数会将cfg.optim_wrapper.type切换为AmpOptimWrapper并设置loss_scale='dynamic'(见 tools/train.py),从而实现自动混合精度(AMP)训练,通常可显著降低显存占用并提升吞吐。

1.2 在 CPU 上训练

模型默认放置在 CUDA 设备上,仅在检测不到 CUDA 设备时才回退到 CPU。因此,如果希望在 CPU 上训练,需要先通过环境变量禁用 GPU 可见性:

export CUDA_VISIBLE_DEVICES=-1

该机制的更多细节可参见 MMEngine 的 Runner 实现(runner.py 中关于设备分配的判断逻辑)。

在 CPU 上训练 MOT 模型 QDTrack 的示例:

CUDA_VISIBLE_DEVICES=-1 python tools/train.py configs/qdtrack/qdtrack_faster-rcnn_r50_fpn_8xb2-4e_mot17halftrain_test-mot17halfval.py

tools/train.py的第一个位置参数config是必填的训练配置文件路径(见 tools/train.py),其余参数均为可选。

1.3 在单 GPU 上训练

如果只有一个 GPU,可以直接使用tools/train.py

python tools/train.py ${CONFIG_FILE} [optional arguments]

可以通过export CUDA_VISIBLE_DEVICES=$GPU_ID选择具体使用哪张 GPU。

在单 GPU 上训练 MOT 模型 QDTrack 的示例(使用 2 号 GPU):

CUDA_VISIBLE_DEVICES=2 python tools/train.py configs/qdtrack/qdtrack_faster-rcnn_r50_fpn_8xb2-4e_mot17halftrain_test-mot17halfval.py

1.4 在单节点多 GPU 上训练

仓库提供了 tools/dist_train.sh 用于在多个 GPU 上启动训练,其基本用法为:

bash ./tools/dist_train.sh ${CONFIG_FILE} ${GPU_NUM} [optional arguments]

从脚本源码可见,它本质上是调用 PyTorch 的torch.distributed.launch,设置--nproc_per_node=$GPUS--master_port=$PORT等分布式参数,并将--launcher pytorch传给训练入口(见 tools/dist_train.sh)。

多作业并行时的端口管理

如果希望在一台机器上同时启动多个作业,例如在拥有 8 个 GPU 的机器上启动 2 个 4-GPU 训练作业,必须为每个作业指定不同的端口(默认端口为 29500),以避免通信冲突:

CUDA_VISIBLE_DEVICES=0,1,2,3 PORT=29500 ./tools/dist_train.sh ${CONFIG_FILE} 4 CUDA_VISIBLE_DEVICES=4,5,6,7 PORT=29501 ./tools/dist_train.sh ${CONFIG_FILE} 4

在单节点多 GPU 上训练 MOT 模型 QDTrack 的示例(8 GPU):

bash ./tools/dist_train.sh configs/qdtrack/qdtrack_faster-rcnn_r50_fpn_8xb2-4e_mot17halftrain_test-mot17halfval.py 8

1.5 在多个节点上训练

当使用以太网连接多台机器时,需要在每台机器上分别运行命令,并通过环境变量协调各节点的角色:

第一台机器(主节点):

NNODES=2 NODE_RANK=0 PORT=$MASTER_PORT MASTER_ADDR=$MASTER_ADDR bash tools/dist_train.sh $CONFIG $GPUS

第二台机器:

NNODES=2 NODE_RANK=1 PORT=$MASTER_PORT MASTER_ADDR=$MASTER_ADDR bash tools/dist_train.sh $CONFIG $GPUS

其中MASTER_ADDR应指向第一台机器的 IP,MASTER_PORT为约定好的通信端口。脚本中NNODESNODE_RANKPORTMASTER_ADDR均有默认值(NNODES=1NODE_RANK=0PORT=29500MASTER_ADDR=127.0.0.1),多节点场景下必须显式覆盖(见 tools/dist_train.sh)。

注意:如果节点之间没有 InfiniBand 等高速网络,多机通信通常会很慢,训练效率将显著下降。

1.6 使用 Slurm 进行训练

Slurm 是计算集群常用的作业调度系统。在 Slurm 管理的集群上,可以使用 tools/slurm_train.sh 提交训练作业,它同时支持单节点和多节点训练。

基本用法:

bash ./tools/slurm_train.sh ${PARTITION} ${JOB_NAME} ${CONFIG_FILE} ${WORK_DIR} ${GPUS}

使用 Slurm 训练 MOT 模型 QDTrack 的完整示例:

PORT=29501 \ GPUS_PER_NODE=8 \ SRUN_ARGS="--quotatype=reserved" \ bash ./tools/slurm_train.sh \ mypartition \ mottrack configs/qdtrack/qdtrack_faster-rcnn_r50_fpn_8xb2-4e_mot17halftrain_test-mot17halfval.py ./work_dirs/QDTrack \ 8

从 tools/slurm_train.sh 源码可以看到脚本内部通过srun提交作业,关键参数包括:

参数说明默认值
PARTITION作业提交到的 Slurm 分区必填(位置参数)
JOB_NAME作业名称必填(位置参数)
GPUS总 GPU 数(即--ntasks8
GPUS_PER_NODE每节点 GPU 数8
CPUS_PER_TASK每个任务分配的 CPU 数5
SRUN_ARGS附加的srun参数空字符串
PORT分布式通信端口29500

脚本最后会以--work-dir=${WORK_DIR} --launcher="slurm"调用 tools/train.py,并将位置参数第 5 个之后的内容(${@:5})作为额外参数透传。

1.7 QDTrack 训练配置解读

以 configs/qdtrack/qdtrack_faster-rcnn_r50_fpn_8xb2-4e_mot17halftrain_test-mot17halfval.py 为例,它通过_base_继承了两部分内容:

_base_ = [ './qdtrack_faster-rcnn_r50_fpn_4e_base.py', '../_base_/datasets/mot_challenge.py', ]

其中 configs/base/datasets/mot_challenge.py 定义了 MOT17 数据集的加载方式:data_root = 'data/MOT17/'、图像尺寸img_scale = (1088, 1088),并使用UniformRefFrameSample采样参考帧、通过TransformBroadcaster对关键帧与参考帧施加相同的随机缩放/裁剪/翻转等数据增强,最后用PackTrackInputs打包(见 mot_challenge.py)。

QDTrack 配置还设置了双重评估器:

val_evaluator = [ dict(type='CocoVideoMetric', metric=['bbox'], classwise=True), dict(type='MOTChallengeMetric', metric=['HOTA', 'CLEAR', 'Identity']) ]

即同时用CocoVideoMetric评估检测质量(bbox AP)与MOTChallengeMetric评估跟踪质量(HOTA / CLEAR / Identity),并设置randomness = dict(seed=6)固定随机种子以保证结果可复现。

二、测试现有跟踪模型

本节介绍如何在受支持的数据集上测试现有跟踪模型,支持以下测试环境:

  • CPU
  • 单 GPU
  • 单节点多 GPU
  • 多节点

同样可以使用 Slurm 管理测试作业。

2.1 测试前须知(重要约定)

权重加载的三条途径:--checkpoint/--detector/--reid

在 MOT 任务中,不同算法的权重加载方式不同:

  • 需要分别加载 ReID 与检测器权重DeepSORTSORTStrongSORT等算法,使用--detector--reid两个参数分别加载;
  • 不需要 ReID 权重ByteTrackOCSORTQDTrack等算法,使用--checkpoint加载即可。

从 tools/test_tracking.py 的源码实现可以看到:--checkpoint通过cfg.load_from注入配置并整体加载,而--detector--reid则分别调用load_checkpoint(model.detector, args.detector)load_checkpoint(model.reid, args.reid)直接加载到模型子模块上。此外脚本做了互斥校验:--checkpoint不能与--detector/--reid同时出现,否则抛出AssertionError

两种测试模式:基于视频与基于图像

仓库提供了两种评估/测试方式:

  • 基于视频的测试:将整个视频序列送入模型,适合 StrongSORT、Mask2Former 等算法;
  • 基于图像的测试:逐帧独立推理,显存占用更小。

有些算法(如 StrongSORT、Mask2Former)只支持基于视频的测试。如果 GPU 内存无法容纳整个视频,可以通过切换采样器类型来改用基于图像的测试:

# 基于视频的测试 sampler=dict(type='DefaultSampler', shuffle=False, round_up=False) # 基于图像的测试 sampler=dict(type='TrackImgSampler')

TrackImgSampler是 MMDetection 为跟踪任务实现的图像级采样器(源码见 track_img_sampler.py),其设计动机从源码注释中可以直接看到:

  • 在测试模式下,默认 PyTorch 采样器一次会输出一整段视频,整段视频送入数据管线将消耗大量显存(MOTChallenge17 数据集上通常需要 ≥20G 显存),而TrackImgSampler保证每次只向数据管线送入一张图像;
  • 在训练模式下,它保证一个 epoch 内视频中的每张图像恰好被随机采样一次。

在 configs/base/datasets/mot_challenge.py 中可以看到,MOT17 的 val/test dataloader 默认使用TrackImgSampler(图像级采样),注释中也给出了切换为视频级采样的方法。

结果保存路径:outfile_prefix

可以通过修改 evaluator 中的outfile_prefix关键字来设置结果保存路径,例如:

val_evaluator = dict(outfile_prefix='results/sort_mot17')

如果未设置outfile_prefix,评估过程会创建一个临时文件,评估结束后自动删除。

仅格式化不评估:format_only=True

如果只需要格式化输出结果而不进行指标评估,可以设置format_only=True,例如:

test_evaluator = dict(type='MOTChallengeMetric', metric=['HOTA', 'CLEAR', 'Identity'], outfile_prefix='sort_mot17_results', format_only=True)

该参数在 MOTChallengeMetric 中有明确定义:format_only=True时仅将结果格式化为官方提交格式并保存,不执行评估计算。

2.2 在 CPU 上测试

与训练一致,模型默认在 CUDA 上运行,只有在没有 CUDA 设备时才回退到 CPU。在 CPU 上测试需要先禁用 GPU 可见性:

CUDA_VISIBLE_DEVICES=-1 python tools/test_tracking.py ${CONFIG_FILE} [optional arguments]

在 CPU 上测试 MOT 模型 SORT 的示例:

CUDA_VISIBLE_DEVICES=-1 python tools/test_tracking.py configs/sort/sort_faster-rcnn_r50_fpn_8xb2-4e_mot17halftrain_test-mot17halfval.py --detector ${CHECKPOINT_FILE}

注意这里 SORT 需要的是检测器权重,因此使用--detector而非--checkpoint

2.3 在单 GPU 上测试

在单 GPU 上测试,可以直接使用tools/test_tracking.py

python tools/test_tracking.py ${CONFIG_FILE} [optional arguments]

可以通过export CUDA_VISIBLE_DEVICES=$GPU_ID选择 GPU。

在单 GPU 上测试 MOT 模型 QDTrack 的示例:

CUDA_VISIBLE_DEVICES=2 python tools/test_tracking.py configs/qdtrack/qdtrack_faster-rcnn_r50_fpn_8xb2-4e_mot17halftrain_test-mot17halfval.py --detector ${CHECKPOINT_FILE}

(若使用 QDTrack 完整权重,也可以直接改用--checkpoint ${CHECKPOINT_FILE}。)

2.4 在单节点多 GPU 上测试

仓库提供了 tools/dist_test_tracking.sh 用于在多个 GPU 上启动测试,基本用法:

bash ./tools/dist_test_tracking.sh ${CONFIG_FILE} ${GPU_NUM} [optional arguments]

在单节点多 GPU 上测试 MOT 模型 DeepSORT 的示例:

bash ./tools/dist_test_tracking.sh configs/qdtrack/qdtrack_faster-rcnn_r50_fpn_8xb2-4e_mot17halftrain_test-mot17halfval.py 8 --detector ${CHECKPOINT_FILE} --reid ${CHECKPOINT_FILE}

示例同时传入--detector--reid,这正是 DeepSORT 这类「检测器 + ReID 模型」两级架构的典型用法。与训练脚本相同,tools/dist_test_tracking.sh 内部通过torch.distributed.launch启动多进程,并将--launcher pytorch传给test_tracking.py

2.5 在多个节点上测试

多节点测试与「多节点训练」的流程类似:在每台机器上设置NNODESNODE_RANKPORTMASTER_ADDR环境变量后运行dist_test_tracking.sh即可,这里不再重复。

2.6 使用 Slurm 进行测试

在 Slurm 管理的集群上,可以使用 tools/slurm_test_tracking.sh 提交测试作业,支持单节点与多节点测试。

基本用法:

[GPUS=${GPUS}] bash tools/slurm_test_tracking.sh ${PARTITION} ${JOB_NAME} ${CONFIG_FILE} [optional arguments]

与训练版本相比,测试脚本没有 WORK_DIR 位置参数,GPUS通过环境变量传入(默认 8),第 4 个位置参数之后的内容(${@:4})全部作为额外参数透传给 tools/test_tracking.py(见 slurm_test_tracking.sh)。

使用 Slurm 测试 VIS 模型 Mask2Former 的示例:

GPUS=8 bash tools/slurm_test_tracking.sh \ mypartition \ vis \ configs/mask2former_vis/mask2former_r50_8xb2-8e_youtubevis2021.py \ --checkpoint ${CHECKPOINT_FILE}

该示例展示了 VIS(视频实例分割)任务如何使用跟踪测试管线:Mask2Former 是典型的「只支持基于视频测试」的模型,其配置位于 configs/mask2former_vis/(含 YouTube-VIS 2019/2021 多套配置),测试时通过--checkpoint加载完整权重。

三、跟踪评估指标与采样机制的原理解读

3.1 MOTChallengeMetric:HOTA / CLEAR / Identity

跟踪评估的默认指标由 MOTChallengeMetric 提供,支持HOTACLEARIdentity三类指标,默认全部启用:

  • HOTA(Higher Order Tracking Accuracy):衡量检测、关联与定位的综合精度;
  • CLEAR:包含 MOTA(多目标跟踪准确率)等经典指标;
  • Identity:包含 IDF1 等身份保持相关指标。

该 metric 的构造函数还支持track_iou_thr(跟踪评估的 IoU 阈值,默认 0.5)、benchmark(评估基准,可选MOT15/MOT16/MOT17/MOT20/DanceTrack,默认MOT17)、postprocess_tracklet_cfg(轨迹后处理,如InterpolateTracklets轨迹插值,支持min_num_framesmax_num_framesuse_gsismooth_tau等参数)以及前文提到的outfile_prefixformat_only

3.2 数据采样的取舍:何时用视频级,何时用图像级

结合前文与源码(track_img_sampler.py),可以给出如下实践建议:

场景推荐采样器原因
显存充足、算法依赖时序信息(StrongSORT、Mask2Former VIS)DefaultSampler(shuffle=False, round_up=False)(视频级)需要整段视频的时序上下文做关联/掩码传播
显存有限、或希望逐帧独立推理TrackImgSampler(图像级)每次仅处理一帧,MOT17 全视频评估可节省大量显存

训练阶段的 QDTrack 等算法在 mot_challenge.py 中默认使用TrackImgSampler,保证一帧一采样且整段视频每帧在一个 epoch 内恰被采样一次。

四、快速参考:命令速查表

场景命令
CPU 训练 QDTrackCUDA_VISIBLE_DEVICES=-1 python tools/train.py configs/qdtrack/qdtrack_faster-rcnn_r50_fpn_8xb2-4e_mot17halftrain_test-mot17halfval.py
单 GPU 训练CUDA_VISIBLE_DEVICES=2 python tools/train.py ${CONFIG_FILE}
单节点多 GPU 训练bash ./tools/dist_train.sh ${CONFIG_FILE} 8
单节点双作业训练CUDA_VISIBLE_DEVICES=0,1,2,3 PORT=29500 ./tools/dist_train.sh ${CONFIG_FILE} 4(另一作业用PORT=29501CUDA_VISIBLE_DEVICES=4,5,6,7
多节点训练NNODES=2 NODE_RANK=0 PORT=$MASTER_PORT MASTER_ADDR=$MASTER_ADDR bash tools/dist_train.sh $CONFIG $GPUS
Slurm 训练PORT=29501 GPUS_PER_NODE=8 SRUN_ARGS="--quotatype=reserved" bash ./tools/slurm_train.sh mypartition mottrack ${CONFIG_FILE} ./work_dirs/QDTrack 8
CPU 测试 SORTCUDA_VISIBLE_DEVICES=-1 python tools/test_tracking.py configs/sort/sort_faster-rcnn_r50_fpn_8xb2-4e_mot17halftrain_test-mot17halfval.py --detector ${CHECKPOINT_FILE}
单 GPU 测试 QDTrackCUDA_VISIBLE_DEVICES=2 python tools/test_tracking.py ${CONFIG_FILE} --checkpoint ${CHECKPOINT_FILE}
单节点多 GPU 测试 DeepSORTbash ./tools/dist_test_tracking.sh ${CONFIG_FILE} 8 --detector ${CHECKPOINT_FILE} --reid ${CHECKPOINT_FILE}
Slurm 测试 Mask2Former VISGPUS=8 bash tools/slurm_test_tracking.sh mypartition vis configs/mask2former_vis/mask2former_r50_8xb2-8e_youtubevis2021.py --checkpoint ${CHECKPOINT_FILE}
混合精度训练在任意训练命令后追加--amp
自动学习率缩放在任意训练命令后追加--auto-scale-lr(需配置含auto_scale_lr字段)

五、相关资源索引

  • 训练入口:tools/train.py
  • 跟踪测试入口:tools/test_tracking.py
  • 分布式训练/测试脚本:tools/dist_train.sh、tools/dist_test_tracking.sh
  • Slurm 训练/测试脚本:tools/slurm_train.sh、tools/slurm_test_tracking.sh
  • 示例配置:configs/qdtrack/、configs/sort/、configs/mask2former_vis/、configs/base/datasets/mot_challenge.py
  • 采样器实现:track_img_sampler.py
  • 评估指标实现:mot_challenge_metric.py、coco_video_metric.py
  • 在线演示脚本(MOT 推理):demo/mot_demo.py
  • 人工智能
  • 计算机视觉
  • 深度学习
  • 模型评测

【免费下载链接】mmdetection

OpenMMLab Detection Toolbox and Benchmark

项目地址:https://gitcode.com/gh_mirrors/mm/mmdetection
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询