- 人工智能
- 计算机视觉
- 深度学习
- 模型评测
【免费下载链接】mmdetection
OpenMMLab Detection Toolbox and Benchmark
本文档是 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.enable与auto_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.pytools/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.py1.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 81.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为约定好的通信端口。脚本中NNODES、NODE_RANK、PORT、MASTER_ADDR均有默认值(NNODES=1、NODE_RANK=0、PORT=29500、MASTER_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 数(即--ntasks) | 8 |
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 与检测器权重:
DeepSORT、SORT、StrongSORT等算法,使用--detector和--reid两个参数分别加载; - 不需要 ReID 权重:
ByteTrack、OCSORT、QDTrack等算法,使用--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 在多个节点上测试
多节点测试与「多节点训练」的流程类似:在每台机器上设置NNODES、NODE_RANK、PORT、MASTER_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 提供,支持HOTA、CLEAR、Identity三类指标,默认全部启用:
- 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_frames、max_num_frames、use_gsi、smooth_tau等参数)以及前文提到的outfile_prefix和format_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 训练 QDTrack | CUDA_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=29501与CUDA_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 测试 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} |
| 单 GPU 测试 QDTrack | CUDA_VISIBLE_DEVICES=2 python tools/test_tracking.py ${CONFIG_FILE} --checkpoint ${CHECKPOINT_FILE} |
| 单节点多 GPU 测试 DeepSORT | bash ./tools/dist_test_tracking.sh ${CONFIG_FILE} 8 --detector ${CHECKPOINT_FILE} --reid ${CHECKPOINT_FILE} |
| Slurm 测试 Mask2Former VIS | GPUS=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
相关推荐
mmdetection 中的 DeepSORT 多目标跟踪:配置解析、模型训练与 MOT17 评测实战
mmdetection 中的 DeepSORT 多目标跟踪:配置解析、模型训练与 MOT17 评测实战 导读 DeepSORT(Simple Online an
人工智能计算机视觉深度学习模型评测MMDetection 视频多目标跟踪(MOT/VIS)训练与测试完全指南:从 CPU 单卡到多机 Slurm 的工程实践
MMDetection 视频多目标跟踪(MOT/VIS)训练与测试完全指南:从 CPU 单卡到多机 Slurm 的工程实践 导读 本文基于 MMDetectio
人工智能计算机视觉深度学习模型评测mmdetection 中的 QDTrack 多目标跟踪实现:Quasi-Dense 相似度学习原理、配置与训练实战
mmdetection 中的 QDTrack 多目标跟踪实现:Quasi Dense 相似度学习原理、配置与训练实战 导读 本文以 mmdetection 仓库
人工智能计算机视觉深度学习模型评测
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考