ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

Libra R-CNN 平衡学习框架解析:MMDetection 中的 IoU 平衡采样、平衡特征金字塔与平衡 L1 损失

Libra R-CNN 平衡学习框架解析:MMDetection 中的 IoU 平衡采样、平衡特征金字塔与平衡 L1 损失 Libra R-CNN 平衡学习框架解析MMDetection 中的 IoU 平衡采样、平衡特征金字塔与平衡 L1 损失【免费下载链接】mmdetectionOpenMMLab Detection Toolbox and Benchmark项目地址: https://gitcode.com/gh_mirrors/mm/mmdetectionLibra R-CNN 是 CVPR 2019 提出的一种面向目标检测的平衡学习框架其核心思想是检测器的性能瓶颈不仅来自网络结构更来自训练过程中样本、特征与目标三个层面的不平衡。本文以 MMDetection 仓库中configs/libra_rcnn/的完整配置与mmdet/models/下的源码实现为据系统拆解 Libra R-CNN 的三大核心组件——IoU-balanced SamplingIoU 平衡采样、Balanced Feature Pyramid平衡特征金字塔BFP与 Balanced L1 Loss平衡 L1 损失并给出各骨干网络变体的完整配置、实验结果与训练/测试实操方法帮助读者在 MMDetection 中直接复现与二次改造 Libra R-CNN。一、背景检测训练中的三重不平衡与网络架构相比同样决定检测器成败的训练过程在很长一段时间内受到的关注较少。Libra R-CNN 重新审视了检测器的标准训练流程发现检测性能常常被训练过程中的不平衡所限制这种不平衡主要存在于三个层面层面不平衡表现Libra R-CNN 的应对组件样本层面Sample level困难样本与简单样本、不同实例间的正样本数量差异悬殊IoU 平衡采样IoU-balanced sampling特征层面Feature level不同层级的特征语义/分辨率不一致高层特征与低层特征信息不均衡平衡特征金字塔Balanced Feature Pyramid目标层面Objective level分类与回归损失梯度量级失衡回归大误差主导梯度平衡 L1 损失Balanced L1 Loss得益于整体平衡的设计论文报告在不引入任何额外技巧without bells and whistles的情况下Libra R-CNN 相比 FPN Faster R-CNN 在 MS COCO 上 AP 提升 2.5 个点相比 RetinaNet 提升 2.0 个点。论文的 IJCV 扩展版进一步将框架泛化到面向实例识别的平衡学习Towards Balanced Learning for Instance Recognition并在 MS COCO、LVIS 与 Pascal VOC 上验证了整体平衡设计的有效性。二、三大核心组件从配置到源码逐层拆解2.1 样本层面IoU 平衡采样随机采样得到的负样本大多是与真实目标 IoU 极低的简单样本它们对训练贡献有限同时一幅图中不同实例拥有的正样本数量可能差异巨大。Libra R-CNN 从两个方向缓解样本层面的不平衡实例平衡正采样Instance Balanced Pos Sampling保证每个真实实例GT被采样到近似等量的正样本。源码实现位于 mmdet/models/task_modules/samplers/instance_balanced_pos_sampler.py其核心逻辑是_sample_pos先按assign_result.gt_inds找出每个实例对应的正样本计算num_per_gt round(num_expected / num_gts) 1再对每个实例分别采样num_per_gt个正样本从而避免大目标/易检实例垄断正样本配额。IoU 平衡负采样IoU Balanced Neg Sampling源码位于 mmdet/models/task_modules/samplers/iou_balanced_neg_sampler.py。它不再单纯随机抽取负样本而是将负样本按与 GT 的 IoU 划分为num_bins个区间bin在每个区间内均匀采样。其sample_via_interval方法按iou_interval (max_iou - floor_thr) / num_bins划分区间并逐 bin 抽取num_expected / num_bins个样本从而保证高 IoU 的困难负样本也有机会被选中。在配置中两者通过CombinedSampler组合使用实现见 mmdet/models/task_modules/samplers/combined_sampler.py分别指定正、负采样器。同时 RPN 阶段设置了neg_pos_ub5负正样本比例上限与allowed_border-1从 RPN 到 RCNN 全程维持采样平衡。2.2 特征层面平衡特征金字塔BFPBFPBalanced Feature Pyramid在标准 FPN 之后串接一个均衡-精炼-回撒模块其完整实现位于 mmdet/models/necks/bfp.py前向过程分三步对应源码 L83-L109Gather均衡汇集将各层特征统一 resize 到refine_level对应的尺寸后取平均得到单一融合特征bsf。低于refine_level的层用F.adaptive_max_pool2d下采样高于或等于的层用F.interpolatenearest上采样Refine精炼对融合后的特征做精炼refine_type支持None不精炼、conv3×3 卷积与non_localNonLocal2dreduction1、use_scaleFalse三种配置中统一采用non_local以捕获全局依赖Scatter残差回撒将精炼结果按各层原尺寸回撒并与原输入做残差相加residual inputs[i]既保留了原金字塔的逐层信息又引入了跨层均衡后的整体特征。BFP 的构造函数参数与配置字段一一对应in_channels各层输入通道须一致通常为 256、num_levels金字塔层数配置为 5、refine_level汇集与精炼层索引自底向上计数Faster R-CNN 取 2、RetinaNet 取 1、refine_type精炼算子类型。源码中assert 0 self.refine_level self.num_levels限定了索引合法性。2.3 目标层面平衡 L1 损失Balanced L1 LossBalanced L1 Loss 旨在使分类与回归任务的梯度量级协调同时抑制回归中离群大误差对梯度主导。实现位于 mmdet/models/losses/balanced_l1_loss.py其核心公式源码 L44-L49记diff |pred - target|b e^(gamma/alpha) - 1当diff beta时loss (alpha / b) * (b * diff 1) * log(b * diff / beta 1) - alpha * diff小误差段梯度被放大训练更充分当diff beta时loss gamma * diff gamma / b - alpha * beta大误差段梯度被截断为常数gamma防止离群点主导训练。三个关键参数含义如下类定义默认值与配置中实际取值一致参数语义默认值Faster R-CNN 配置RetinaNet 配置alpha小误差段的梯度放大系数分母0.50.50.5gamma大误差段的梯度截断常数1.51.51.5beta分段函数的分界阈值预测与目标差值的分界点1.01.00.11注意beta在 Faster R-CNN 与 RetinaNet 配置中取值不同RetinaNet 使用更小的beta0.11这是因为密集检测头dense head的回归目标分布与两阶段 RCNN 头不同需要更早进入梯度截断区。Loss 本身由weighted_loss装饰器包装支持none/mean/sum三种 reduction并乘以loss_weight1.0。三、配置文件逐行解析以 Faster R-CNN 变体为例Libra R-CNN 的全部实验配置集中在 configs/libra_rcnn/ 目录共有 5 个配置文件与 1 个模型元信息文件metafile.yml。所有配置都通过继承基础检测器配置实现最小化改动。libra-faster-rcnn_r50_fpn_1x_coco.py 继承自 configs/faster_rcnn/faster-rcnn_r50_fpn_1x_coco.py完整改动如下_base_ ../faster_rcnn/faster-rcnn_r50_fpn_1x_coco.py # model settings model dict( neck[ dict( typeFPN, in_channels[256, 512, 1024, 2048], out_channels256, num_outs5), dict( typeBFP, in_channels256, num_levels5, refine_level2, refine_typenon_local) ], roi_headdict( bbox_headdict( loss_bboxdict( _delete_True, typeBalancedL1Loss, alpha0.5, gamma1.5, beta1.0, loss_weight1.0))), # model training and testing settings train_cfgdict( rpndict(samplerdict(neg_pos_ub5), allowed_border-1), rcnndict( samplerdict( _delete_True, typeCombinedSampler, num512, pos_fraction0.25, add_gt_as_proposalsTrue, pos_samplerdict(typeInstanceBalancedPosSampler), neg_samplerdict( typeIoUBalancedNegSampler, floor_thr-1, floor_fraction0, num_bins3)))))配置要点说明neck 改为列表FPN 与 BFP 串接FPN 输出 256 通道的 5 层特征BFP 在其后做均衡精炼_delete_TrueMMEngine 配置继承中用于删除基类同名键。这里用于1替换基类bbox_head.loss_bbox为BalancedL1Loss2替换基类train_cfg.rcnn.sampler为CombinedSampler避免与基类默认的随机采样器混叠采样器参数RCNN 阶段每个 batch 采样num512个样本、正样本比例pos_fraction0.25、add_gt_as_proposalsTrue把 GT 也纳入候选负采样器floor_thr-1表示全部负样本都走 IoU 平衡采样不区分地板区floor_fraction0表示地板区样本占比为 0num_bins3将负样本按 IoU 分 3 个区间RPN 阶段neg_pos_ub5限制负样本不超过正样本的 5 倍allowed_border-1允许预测框完全位于图像外不裁剪到图像边界对应论文中扩边实验设置。3.1 骨干网络变体与 RetinaNet 变体libra-faster-rcnn_r101_fpn_1x_coco.py 继承 R50 版配置仅替换backbone.depth101并加载torchvision://resnet101预训练权重libra-faster-rcnn_x101-64x4d_fpn_1x_coco.py 继承 R50 版配置将骨干替换为 ResNeXt-101-64x4dgroups64、base_width4预训练权重为open-mmlab://resnext101_64x4dlibra-retinanet_r50_fpn_1x_coco.py 继承 configs/retinanet/retinanet_r50_fpn_1x_coco.pyneck 中的 FPN 使用单阶段检测器专属配置start_level1、add_extra_convson_inputBFP 的refine_level1并将bbox_head.loss_bbox替换为beta0.11的BalancedL1Loss无需_delete_因为直接覆盖了loss_bbox键对应的字段值。3.2 Fast R-CNN 变体与预生成 Proposallibra-fast-rcnn_r50_fpn_1x_coco.py 继承 configs/fast_rcnn/fast-rcnn_r50_fpn_1x_coco.py同样挂载 BFP、CombinedSampler与BalancedL1Loss区别在于 Fast R-CNN 不训练 RPN而是直接读取离线生成的候选框文件train_dataloader dict( datasetdict(proposal_filelibra_proposals/rpn_r50_fpn_1x_train2017.pkl)) val_dataloader dict( datasetdict(proposal_filelibra_proposals/rpn_r50_fpn_1x_val2017.pkl)) test_dataloader val_dataloader该配置文件中还以注释形式给出使用_base_字段原地修改的等价写法_base_.train_dataloader.dataset.proposal_file ...两种方式皆受支持可按习惯选择。README 的结果表中该变体未给出复现指标属于需要自行准备libra_proposals/*.pkl预生成候选文件后才能训练/评估的实验性配置。四、实验结果与模型对比COCO 2017 val下表为 README 中报告的 COCO 2017 val 上的结果test-dev 指标通常略高于 val。表中的推理速度fps在 V100 上以 batch size 1、FP32、分辨率 (800, 1333) 测得训练资源为 8× V100学习率调度为 1x12 epochs详见 metafile.yml架构骨干风格Lr schd显存 (GB)推理速度 (fps)box AP配置文件Faster R-CNNR-50-FPNpytorch1x4.619.038.3libra-faster-rcnn_r50_fpn_1x_coco.pyFast R-CNNR-50-FPNpytorch1x———libra-fast-rcnn_r50_fpn_1x_coco.pyFaster R-CNNR-101-FPNpytorch1x6.514.440.1libra-faster-rcnn_r101_fpn_1x_coco.pyFaster R-CNNX-101-64x4d-FPNpytorch1x10.88.542.7libra-faster-rcnn_x101-64x4d_fpn_1x_coco.pyRetinaNetR-50-FPNpytorch1x4.217.737.6libra-retinanet_r50_fpn_1x_coco.py从横向对比看在同等 Faster R-CNN R-50-FPN 条件下Libra 版本 38.3 AP 高于标准 FPN Faster R-CNNRetinaNet 变体 37.6 AP 也高于论文报告的标准 RetinaNet。各模型权重与训练日志可通过 MMDetection 模型库Model Zoo按上述配置名检索下载metafile.yml中登记了每个模型的权重地址与推理耗时。五、训练、测试与引用5.1 训练与测试在安装好依赖并完成数据集准备后使用 MMDetection 标准入口脚本即可训练与评测# 单卡训练 Faster R-CNN R-50-FPN 变体 python tools/train.py configs/libra_rcnn/libra-faster-rcnn_r50_fpn_1x_coco.py # 多卡分布式训练 bash tools/dist_train.sh configs/libra_rcnn/libra-faster-rcnn_r50_fpn_1x_coco.py 8 # 测试评估 python tools/test.py configs/libra_rcnn/libra-faster-rcnn_r50_fpn_1x_coco.py checkpoint路径 --eval bbox需要注意**Fast R-CNN 变体libra-fast-rcnn_r50_fpn_1x_coco.py**依赖libra_proposals/下预生成的 RPN 候选框文件rpn_r50_fpn_1x_train2017.pkl/rpn_r50_fpn_1x_val2017.pkl运行前需确保这些文件存在且路径与配置一致若使用 Slurm 集群可参考 tools/slurm_train.sh 与 tools/slurm_test.sh。5.2 论文引用仓库提供的配置用于复现 CVPR 2019 论文 Libra R-CNN 的实验结果其 IJCV 扩展版也已公开发表引用信息如下inproceedings{pang2019libra, title{Libra R-CNN: Towards Balanced Learning for Object Detection}, author{Pang, Jiangmiao and Chen, Kai and Shi, Jianping and Feng, Huajun and Ouyang, Wanli and Dahua Lin}, booktitle{IEEE Conference on Computer Vision and Pattern Recognition}, year{2019} } article{pang2021towards, title{Towards Balanced Learning for Instance Recognition}, author{Pang, Jiangmiao and Chen, Kai and Li, Qi and Xu, Zhihai and Feng, Huajun and Shi, Jianping and Ouyang, Wanli and Lin, Dahua}, journal{International Journal of Computer Vision}, volume{129}, number{5}, pages{1376--1393}, year{2021}, publisher{Springer} }六、扩展阅读建议想深入理解 BFP 的 gather-refine-scatter 三阶段细节直接阅读 mmdet/models/necks/bfp.py 的forward实现L79-L111想调整损失曲线形状阅读 mmdet/models/losses/balanced_l1_loss.py 中balanced_l1_loss的分段公式再按需修改alpha/gamma/beta想复刻自定义采样策略可分别参考InstanceBalancedPosSampler、IoUBalancedNegSampler与CombinedSampler三个采样器的源码理解AssignResult与采样器接口的协作方式平衡思想同样体现在 MMDetection 其他算法中例如sablSABL 采用 Bucket 回归解耦分类/回归等目录可与 Libra R-CNN 对照学习。总体而言Libra R-CNN 用三个轻量组件分别命中训练流程中的样本、特征与目标三处不平衡MMDetection 以继承基配置 最小改动的方式将其完整落地是理解训练策略也是检测性能上限的一部分这一思想的最佳入门样例之一。【免费下载链接】mmdetectionOpenMMLab Detection Toolbox and Benchmark项目地址: https://gitcode.com/gh_mirrors/mm/mmdetection创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进