CentripetalNet 在 MMDetection 中的实现与实战指南:基于向心偏移的高质量角点配对检测器

发布时间:2026/9/19 1:57:11
CentripetalNet 在 MMDetection 中的实现与实战指南:基于向心偏移的高质量角点配对检测器 CentripetalNet 在 MMDetection 中的实现与实战指南基于向心偏移的高质量角点配对检测器【免费下载链接】mmdetectionOpenMMLab Detection Toolbox and Benchmark项目地址: https://gitcode.com/gh_mirrors/mm/mmdetectionCentripetalNet 是 OpenMMLab MMDetection 工具箱中内置的基于关键点Keypoint的 Anchor-Free 目标检测器实现。它以向心偏移Centripetal Shift替代 CornerNet 的嵌入向量Embedding来配对同一实例的左上/右下角点并引入 Cross-star Deformable Convolution 做特征自适应显著提升角点配对的准确率。本文将围绕 configs/centripetalnet/README.md 展开结合仓库源码、完整配置与测试用例讲解该算法在 MMDetection 中的结构、配置、训练、测试与 TTA 全流程帮助你直接复现并深入理解这一经典检测器。CentripetalNet 核心思想从嵌入向量到向心偏移关键点检测器Keypoint-based Detector通过预测物体的左上角与右下角关键点来生成边界框其性能往往受限于角点匹配错误——即左上角与右下角点虽然各自预测准确却无法正确配对到同一实例。CornerNet 使用关联嵌入Associative Embedding来解决配对问题但嵌入向量的度量空间学习难度较大容易出现误配。CentripetalNet 提出了一种更直观、信息量更丰富的配对方式为每个角点额外预测一个向心偏移Centripetal Shift。如 configs/centripetalnet/README.md 的 Abstract 所述CentripetalNet predicts the position and the centripetal shift of the corner points and matches corners whose shifted results are aligned. Combining position information, our approach matches corner points more accurately than the conventional embedding approaches do.具体而言对左上角点向心偏移指向物体中心的方向向右下对右下角点向心偏移指向物体中心的方向向左上若两个角点属于同一实例它们各自向心移动后的位置应在物体中心附近重合据此即可完成配对。该方法相比传统嵌入方法能更精确地匹配角点。同时为了增强角点处对框内信息的感知作者设计了Cross-star Deformable Convolution进行特征自适应feature adaption。论文还探索了在 Anchor-Free 检测器上通过掩码预测模块扩展实例分割任务在 MS-COCO test-dev 上取得了 48.0% AP 的检测精度并可与当时先进的实例分割方法40.2% MaskAP相媲美引自该 README 的 Abstract。已发布模型与复现指标README 的 Results and Models 表格给出了官方发布模型的复现信息仅一个配置BackboneBatch SizeStep/Total EpochsMem (GB)Inf time (fps)box APConfigDownloadHourglassNet-10416 x 6190/21016.73.744.8configmodel | log该表提供了两个重要复现注意事项README Note 部分TTA 设置官方 TTATest-Time Augmentation为单尺度 flipTrue。要复现 TTA 精度需在测试命令中追加--tta参数。检查点选择官方发布的是最优 checkpoint而非最后一个 checkpointbox AP 44.8 vs 训练末期 44.6。模型结构与配置逐项解析核心配置文件为 configs/centripetalnet/centripetalnet_hourglass104_16xb6-crop511-210e-mstest_coco.py它基于default_runtime与 COCO 检测数据集配置组合而成。数据预处理器data_preprocessor dict( typeDetDataPreprocessor, mean[123.675, 116.28, 103.53], std[58.395, 57.12, 57.375], bgr_to_rgbTrue)使用 ImageNet 统计的均值/方差做归一化bgr_to_rgbTrue表示将 BGR 输入转为 RGB。注意后续RandomCenterCropPad中的to_rgbdata_preprocessor[bgr_to_rgb]注释明确说明图像数据不会被转为 RGBImage data is not converted to rgb即填充逻辑与预处理器的通道顺序保持一致。模型主体model dict( typeCornerNet, data_preprocessordata_preprocessor, backbonedict( typeHourglassNet, downsample_times5, num_stacks2, stage_channels[256, 256, 384, 384, 384, 512], stage_blocks[2, 2, 2, 2, 2, 4], norm_cfgdict(typeBN, requires_gradTrue)), neckNone, bbox_headdict( typeCentripetalHead, num_classes80, in_channels256, num_feat_levels2, corner_emb_channels0, loss_heatmapdict( typeGaussianFocalLoss, alpha2.0, gamma4.0, loss_weight1), loss_offsetdict(typeSmoothL1Loss, beta1.0, loss_weight1), loss_guiding_shiftdict( typeSmoothL1Loss, beta1.0, loss_weight0.05), loss_centripetal_shiftdict( typeSmoothL1Loss, beta1.0, loss_weight1)), train_cfgNone, test_cfgdict( corner_topk100, local_maximum_kernel3, distance_threshold0.5, score_thr0.05, max_per_img100, nmsdict(typesoft_nms, iou_threshold0.5, methodgaussian)))各组件的作用如下检测器类型整体使用CornerNet单阶段检测器源码位于 mmdet/models/detectors/cornernet.py该类继承SingleStageDetectorneckNone表示不使用 NeckHourglassNet 直接输出多层特征。BackboneHourglassNet-104downsample_times5num_stacks2双沙漏堆叠第二层输出用于最终预测stage 通道与残差块数量按stage_channels/stage_blocks配置。HeadCentripetalHead实现位于 mmdet/models/dense_heads/centripetal_head.py。其中num_feat_levels2HourglassNet-104 同时输出最终特征与中间监督特征故为 2 个层级源码注释说明 HourglassNet-52 仅输出最终特征此时应为 1corner_emb_channels0关闭嵌入分支这是 CentripetalHead 与 CornerHead 的关键差异——用 guiding shift centripetal shift 替代 embedding 完成配对三个损失角点热图使用GaussianFocalLoss(alpha2.0, gamma4.0)CornerNet 论文中变体 focal loss实现见 mmdet/models/losses/gaussian_focal_loss.py偏移与两种 shift 均使用SmoothL1Loss(beta1.0)。test_cfg 解码参数corner_topk100每张图从热图中取前 100 个角点local_maximum_kernel33×3 局部极大值池化核用于热图 NMSdistance_threshold0.5配对距离阈值见下文解码逻辑score_thr0.05、max_per_img100输出过滤阈值与每图最大框数nms使用soft_nms高斯法iou_threshold0.5。训练与测试数据流水线训练流水线train_pipeline包含train_pipeline [ dict(typeLoadImageFromFile, backend_args_base_.backend_args), dict(typeLoadAnnotations, with_bboxTrue), dict( typePhotoMetricDistortion, brightness_delta32, contrast_range(0.5, 1.5), saturation_range(0.5, 1.5), hue_delta18), dict( typeRandomCenterCropPad, crop_size(511, 511), ratios(0.6, 0.7, 0.8, 0.9, 1.0, 1.1, 1.2, 1.3), test_modeFalse, test_pad_modeNone, meandata_preprocessor[mean], stddata_preprocessor[std], to_rgbdata_preprocessor[bgr_to_rgb]), dict(typeResize, scale(511, 511), keep_ratioFalse), dict(typeRandomFlip, prob0.5), dict(typePackDetInputs), ]要点PhotoMetricDistortion做颜色抖动增强RandomCenterCropPad是 CornerNet 系算法训练的关键增强随机裁剪中心区域并填充到方形。配置注释特别指出训练中裁剪图像会被填充为正方形但尺寸可能小于 crop_sizeResize统一缩放到 511×511COCO 上 CornerNet 系列的标准输入尺寸RandomFlip(prob0.5)随机翻转。测试流水线test_pipeline与训练不同不做 Resize而是使用RandomCenterCropPad的test_modeTrue并以test_pad_mode[logical_or, 127]将图像逻辑或填充至边界值 127PackDetInputs需要保留border元信息用于解码时还原裁剪偏移。test_pipeline [ dict(typeLoadImageFromFile, to_float32True, backend_args_base_.backend_args), dict( typeRandomCenterCropPad, crop_sizeNone, ratiosNone, borderNone, test_modeTrue, test_pad_mode[logical_or, 127], meandata_preprocessor[mean], stddata_preprocessor[std], to_rgbdata_preprocessor[bgr_to_rgb]), dict(typeLoadAnnotations, with_bboxTrue), dict( typePackDetInputs, meta_keys(img_id, img_path, ori_shape, img_shape, border)) ]数据加载、优化器与学习率调度train_dataloader dict( batch_size6, num_workers3, batch_samplerNone, datasetdict(pipelinetrain_pipeline)) val_dataloader dict(datasetdict(pipelinetest_pipeline)) test_dataloader val_dataloader单卡 batch_size6、3 个 workerbatch_samplerNone表示关闭 MMEngine 默认的 InfiniteSampler 行为优化器使用Adamlr0.0005并配合clip_grad(max_norm35, norm_type2)梯度裁剪HourglassNet 训练不稳定裁剪必不可少optim_wrapper dict( typeOptimWrapper, optimizerdict(typeAdam, lr0.0005), clip_graddict(max_norm35, norm_type2)) max_epochs 210 param_scheduler [ dict(typeLinearLR, start_factor1.0 / 3, by_epochFalse, begin0, end500), dict( typeMultiStepLR, begin0, endmax_epochs, by_epochTrue, milestones[190], gamma0.1) ] train_cfg dict(typeEpochBasedTrainLoop, max_epochsmax_epochs, val_interval1)学习率采用两段式前 500 次迭代从 1/3 倍学习率线性预热LinearLR190 epoch 处按MultiStepLR衰减 0.1 倍总计训练 210 epoch——这与 README 表格中的 190/210 完全对应。自动学习率缩放# base_batch_size (16 GPUs) x (6 samples per GPU) auto_scale_lr dict(base_batch_size96)配置末尾通过auto_scale_lr dict(base_batch_size96)声明基准批量大小16 卡 × 6 样本。当实际训练批量不同时MMEngine 会依据批量大小自动线性缩放学习率。注释明确提示该值由官方基准确定用户不应修改。源码级原理CentripetalHead 的网络结构CentripetalHead继承自 mmdet/models/dense_heads/corner_head.py 中的CornerHead其构造逻辑centripetal_head.py先调用父类初始化再构建两个新分支loss_guiding_shiftguiding shift 损失默认权重 0.05较小因为它仅起引导作用loss_centripetal_shift向心偏移损失默认权重 1。构造函数还断言centripetal_shift_channels 2且guiding_shift_channels 2各对应 x/y 两个方向分量。各分支与 Cross-star Deformable Convolution_init_centripetal_layerscentripetal_head.py为每个特征层级构建了 8 组子模块每组又分为tl_左上与br_右下两部分feat_adaptionDeformConv2d可变形卷积即论文中的 Cross-star Deformable Convolution输入输出均为in_channels卷积核feat_adaption_conv_kernel3guiding_shift预测 2 通道 guiding shift从角点指向中心由 3×3 Conv 1×1 Conv 堆叠而成dcn_offset由 guiding shift 特征预测可变形卷积的偏移量通道数为kernel² × 2 18centripetal_shift预测 2 通道向心偏移。前向流程forward_singlecentripetal_head.py清晰地体现了 Cross-star Deformable Convolution 的自适应思想先通过父类CornerHead.forward_single得到角点热图、偏移以及角点池化特征tl_pool/br_poolCornerNet 的双向 Corner Pooling 模块BiCornerPool定义在 corner_head.py沿 top/left 与 bottom/right 两个方向池化从池化特征预测guiding shiftguiding shiftdetach 后经dcn_offset生成可变形卷积偏移量——detach 的目的是让 guiding shift 分支作为稳定的引导不被 DCN 反向传播干扰用该偏移量对池化特征做可变形卷积特征自适应在自适应后的特征上预测centripetal shift。最终输出 8 个张量tl_heat, br_heat, tl_off, br_off, tl_guiding_shift, br_guiding_shift, tl_centripetal_shift, br_centripetal_shift。目标生成三种监督信号的数学定义CornerHead.get_targetscorner_head.py同时承担 CentripetalNet 的目标生成关键代码如下# Guiding shift is a kind of offset, from center to corner if with_guiding_shift: gt_tl_guiding_shift[batch_id, 0, top_idx, left_idx] scale_center_x - left_idx gt_tl_guiding_shift[batch_id, 1, top_idx, left_idx] scale_center_y - top_idx gt_br_guiding_shift[batch_id, 0, bottom_idx, right_idx] right_idx - scale_center_x gt_br_guiding_shift[batch_id, 1, bottom_idx, right_idx] bottom_idx - scale_center_y # Centripetal shift is also a kind of offset, from center to corner # and normalized by log. if with_centripetal_shift: gt_tl_centripetal_shift[batch_id, 0, top_idx, left_idx] log(scale_center_x - scale_left) gt_tl_centripetal_shift[batch_id, 1, top_idx, left_idx] log(scale_center_y - scale_top) gt_br_centripetal_shift[batch_id, 0, bottom_idx, right_idx] log(scale_right - scale_center_x) gt_br_centripetal_shift[batch_id, 1, bottom_idx, right_idx] log(scale_bottom - scale_center_y)Guiding shift从角点指向中心的普通偏移角点坐标 → 中心坐标Centripetal shift从角点指向中心、但以中心到角点的距离取对数归一化的偏移log(scale_center_x - scale_left)等即源码注释所说 a kind of offset, from center to corner and normalized by log。取对数可对不同尺度目标提供更均衡的回归目标。热图目标沿用 CornerNet 的高斯半径方法通过gaussian_radiusmin_overlap0.3计算半径后用gen_gaussian_target在真实角点位置绘制高斯核偏移目标只存放下采样取整造成的亚像素残差。损失计算loss_by_featcentripetal_head.py组织 4 组损失det_loss两角点热图的 GaussianFocalLossoff_loss角点偏移 SmoothL1Lossguiding_lossguiding shift 的 SmoothL1Losscentripetal_loss向心偏移的 SmoothL1Loss。loss_by_feat_singlecentripetal_head.py的细节值得注意所有 shift 损失都通过gt_*_heatmap.eq(1)生成掩码即只在真实角点位置计算损失tl_mask/br_mask类无关、shape 为 batch×1×H×W避免对背景像素引入无意义回归guiding 与 centripetal 损失分别取左上/右下两分支的平均。解码与配对距离阈值与中心区域约束推理时predict_by_feat与_decode_heatmapcorner_head.py执行关键点解码与配对对热图做 3×3 局部极大值抑制get_local_maximum取每类 top-k 角点对角点坐标加上预测的亚像素偏移对向心偏移取指数还原tl_centripetal_shift.exp()从而得到两个角点各自指向的预测中心(tl_ctxs, tl_ctys)与(br_ctxs, br_ctys)依据论文 4.1 节的魔法数计算中心区域rcentralmu 1/2.4大面积框area_bboxes 3500时用1/2.1并用面积比值计算距离度量dists area_ct_bboxes / area_rcentral通过约束过滤候选框预测中心落在rcentral区域之外则剔除scores[tl_ctx_inds] -1等 4 个方向距离度量超过distance_threshold配置为 0.5则剔除dist_inds类别不一致、宽高非法br_xs tl_xs等的框剔除剩余候选按平均分数 top-k最后经soft_nms_bboxes_nms输出。解码逻辑同时兼容 embedding 与 centripetal shift 两种配对方式并通过断言with_embedding with_centripetal_shift 1强制二选一。代码注释还提到一个工程细节热图 top-k 展开使用repeat而非expand因为expand是浅拷贝会导致测试阶段 mAP 下降约 10%。如何训练与测试在仓库根目录下使用官方脚本即可完成训练数据路径需按 configs/base/datasets/coco_detection.py 中的data_root准备 COCO 格式数据集# 单卡训练 python tools/train.py configs/centripetalnet/centripetalnet_hourglass104_16xb6-crop511-210e-mstest_coco.py # 多卡8 卡分布式训练 bash tools/dist_train.sh configs/centripetalnet/centripetalnet_hourglass104_16xb6-crop511-210e-mstest_coco.py 8 # 测试普通推理 python tools/test.py configs/centripetalnet/centripetalnet_hourglass104_16xb6-crop511-210e-mstest_coco.py checkpoint.pth # 开启 TTA 复现 README 中的精度 python tools/test.py configs/centripetalnet/centripetalnet_hourglass104_16xb6-crop511-210e-mstest_coco.py checkpoint.pth --tta注意单卡直接运行将使用配置的 batch_size6而auto_scale_lr声明的是 96 的基准批量如需严格对齐官方指标建议按 16 卡 × 6 的配置进行分布式训练或使用自动学习率缩放。TTA 配置解析配置尾部定义了tta_model与tta_pipelinetta_model dict( typeDetTTAModel, tta_cfgdict( nmsdict(typesoft_nms, iou_threshold0.5, methodgaussian), max_per_img100))tta_pipeline将测试变换组织为多组候选变换的笛卡尔积RandomFlipprob1 与 prob0 两档实现原始图 水平翻转、RandomCenterCropPadtest_mode、LoadAnnotations与PackDetInputs。配置注释特别提醒了一个易错点RandomFlipmust be placed beforeRandomCenterCropPad, otherwise bounding box coordinates after flipping cannot be recovered correctly.即翻转变换必须位于中心裁剪填充之前否则翻转后的边界框坐标无法正确还原。PackDetInputs的meta_keys因此额外包含了flip与flip_direction。测试用例验证仓库为CentripetalHead提供了独立单测tests/test_models/test_dense_heads/test_centripetal_head.py覆盖两种关键场景空 GT 场景构造num_classes4, in_channels1, corner_emb_channels0的 head输入两个特征层级断言det_loss 0鼓励预测背景而guiding_loss、centripetal_loss、off_loss均为 0无真实框时不应产生回归损失两个 GT 框场景传入两个真实框断言 4 组损失全部大于 0验证训练信号完整生效。该测试同时印证了 head 的输入规格num_feat_levels个特征层级、forward输出可直接喂入loss_by_feat并确认corner_emb_channels0时不会构建 embedding 分支。引用如需引用 CentripetalNet请使用官方提供的 BibTeX见 configs/centripetalnet/README.mdInProceedings{Dong_2020_CVPR, author {Dong, Zhiwei and Li, Guoxuan and Liao, Yue and Wang, Fei and Ren, Pengju and Qian, Chen}, title {CentripetalNet: Pursuing High-Quality Keypoint Pairs for Object Detection}, booktitle {Proceedings of the IEEE/CVF Conference on Computer Vision and Pattern Recognition (CVPR)}, month {June}, year {2020} }小结CentripetalNet 在 MMDetection 中的实现完整覆盖了论文提出的三大技术点向心偏移配对替代嵌入向量、Cross-star Deformable Convolution 特征自适应feat_adaptiondcn_offset以及中心区域约束解码rcentraldistance_threshold。通过 configs/centripetalnet/centripetalnet_hourglass104_16xb6-crop511-210e-mstest_coco.py 一份配置即可复现 44.8 box APCOCO test-dev配合 TTA深入阅读 centripetal_head.py 与 corner_head.py 则能掌握角点类检测器从目标生成、损失计算到解码配对的完整链路为在此基础上改进或迁移到新任务提供可靠参照。【免费下载链接】mmdetectionOpenMMLab Detection Toolbox and Benchmark项目地址: https://gitcode.com/gh_mirrors/mm/mmdetection创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

关于本文作者

来自尧图内容编辑团队

尧图内容编辑团队 内容团队

尧图内容编辑团队

本文由尧图网络内容编辑团队执笔。团队由资深项目经理、前端工程师与设计师组成,所有内容均来自亲手交付的真实项目,先讲清问题、再给出可落地的解法。尧图深耕北京网站建设十年,服务过京华建材集团、智造科技等各行业客户,把一线经验沉淀为可复用的行业观察。

  • 十年建站经验,覆盖建材、制造、服务、文创等
  • 项目经理把关选题与事实准确性
  • 工程师与设计师联合撰写专业细节
  • 统一编辑规范,保证文风与排版一致
  • 每月复盘转化数据,迭代选题方向

延伸阅读

相关资讯与近期热门内容

深度阅读推荐

建站决策前值得细读的三篇

网站改版的5个关键决策
2024-08-12

网站改版的5个关键决策

什么时候该改版、改到什么程度、如何避免流量掉光,京华建材集团改版复盘给出答案。

获取专属建站方案

看完文章,把您的行业与预算告诉我们,免费获取一份量身定制的官网建设方案与报价。

立即免费咨询