Libra R-CNN 平衡学习框架解析:MMDetection 中的 IoU 平衡采样、平衡特征金字塔与平衡 L1 损失
【免费下载链接】mmdetectionOpenMMLab Detection Toolbox and Benchmark项目地址: https://gitcode.com/gh_mirrors/mm/mmdetection
Libra R-CNN 是 CVPR 2019 提出的一种面向目标检测的"平衡学习"框架,其核心思想是:检测器的性能瓶颈不仅来自网络结构,更来自训练过程中样本、特征与目标三个层面的不平衡。本文以 MMDetection 仓库中configs/libra_rcnn/的完整配置与mmdet/models/下的源码实现为据,系统拆解 Libra R-CNN 的三大核心组件——IoU-balanced Sampling(IoU 平衡采样)、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_ub=5(负正样本比例上限)与allowed_border=-1,从 RPN 到 RCNN 全程维持采样平衡。
2.2 特征层面:平衡特征金字塔(BFP)
BFP(Balanced Feature Pyramid)在标准 FPN 之后串接一个"均衡-精炼-回撒"模块,其完整实现位于 mmdet/models/necks/bfp.py,前向过程分三步(对应源码 L83-L109):
- Gather(均衡汇集):将各层特征统一 resize 到
refine_level对应的尺寸后取平均,得到单一融合特征bsf。低于refine_level的层用F.adaptive_max_pool2d下采样,高于或等于的层用F.interpolate(nearest)上采样; - Refine(精炼):对融合后的特征做精炼,
refine_type支持None(不精炼)、conv(3×3 卷积)与non_local(NonLocal2d,reduction=1、use_scale=False)三种,配置中统一采用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 Loss)
Balanced 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.5 | 0.5 | 0.5 |
gamma | 大误差段的梯度截断常数 | 1.5 | 1.5 | 1.5 |
beta | 分段函数的分界阈值(预测与目标差值的分界点) | 1.0 | 1.0 | 0.11 |
注意beta在 Faster R-CNN 与 RetinaNet 配置中取值不同:RetinaNet 使用更小的beta=0.11,这是因为密集检测头(dense head)的回归目标分布与两阶段 RCNN 头不同,需要更早进入梯度截断区。Loss 本身由@weighted_loss装饰器包装,支持none/mean/sum三种 reduction,并乘以loss_weight=1.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( type='FPN', in_channels=[256, 512, 1024, 2048], out_channels=256, num_outs=5), dict( type='BFP', in_channels=256, num_levels=5, refine_level=2, refine_type='non_local') ], roi_head=dict( bbox_head=dict( loss_bbox=dict( _delete_=True, type='BalancedL1Loss', alpha=0.5, gamma=1.5, beta=1.0, loss_weight=1.0))), # model training and testing settings train_cfg=dict( rpn=dict(sampler=dict(neg_pos_ub=5), allowed_border=-1), rcnn=dict( sampler=dict( _delete_=True, type='CombinedSampler', num=512, pos_fraction=0.25, add_gt_as_proposals=True, pos_sampler=dict(type='InstanceBalancedPosSampler'), neg_sampler=dict( type='IoUBalancedNegSampler', floor_thr=-1, floor_fraction=0, num_bins=3)))))配置要点说明:
- neck 改为列表:FPN 与 BFP 串接,FPN 输出 256 通道的 5 层特征,BFP 在其后做均衡精炼;
_delete_=True:MMEngine 配置继承中用于删除基类同名键。这里用于(1)替换基类bbox_head.loss_bbox为BalancedL1Loss;(2)替换基类train_cfg.rcnn.sampler为CombinedSampler,避免与基类默认的随机采样器混叠;- 采样器参数:RCNN 阶段每个 batch 采样
num=512个样本、正样本比例pos_fraction=0.25、add_gt_as_proposals=True把 GT 也纳入候选;负采样器floor_thr=-1表示全部负样本都走 IoU 平衡采样(不区分地板区),floor_fraction=0表示地板区样本占比为 0,num_bins=3将负样本按 IoU 分 3 个区间; - RPN 阶段:
neg_pos_ub=5限制负样本不超过正样本的 5 倍,allowed_border=-1允许预测框完全位于图像外(不裁剪到图像边界,对应论文中"扩边"实验设置)。
3.1 骨干网络变体与 RetinaNet 变体
- libra-faster-rcnn_r101_fpn_1x_coco.py 继承 R50 版配置,仅替换
backbone.depth=101并加载torchvision://resnet101预训练权重; - libra-faster-rcnn_x101-64x4d_fpn_1x_coco.py 继承 R50 版配置,将骨干替换为 ResNeXt-101-64x4d(
groups=64、base_width=4),预训练权重为open-mmlab://resnext101_64x4d; - libra-retinanet_r50_fpn_1x_coco.py 继承 configs/retinanet/retinanet_r50_fpn_1x_coco.py,neck 中的 FPN 使用单阶段检测器专属配置(
start_level=1、add_extra_convs='on_input'),BFP 的refine_level=1,并将bbox_head.loss_bbox替换为beta=0.11的BalancedL1Loss(无需_delete_,因为直接覆盖了loss_bbox键对应的字段值)。
3.2 Fast R-CNN 变体与预生成 Proposal
libra-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( dataset=dict(proposal_file='libra_proposals/rpn_r50_fpn_1x_train2017.pkl')) val_dataloader = dict( dataset=dict(proposal_file='libra_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,学习率调度为 1x(12 epochs),详见 metafile.yml:
| 架构 | 骨干 | 风格 | Lr schd | 显存 (GB) | 推理速度 (fps) | box AP | 配置文件 |
|---|---|---|---|---|---|---|---|
| Faster R-CNN | R-50-FPN | pytorch | 1x | 4.6 | 19.0 | 38.3 | libra-faster-rcnn_r50_fpn_1x_coco.py |
| Fast R-CNN | R-50-FPN | pytorch | 1x | — | — | — | libra-fast-rcnn_r50_fpn_1x_coco.py |
| Faster R-CNN | R-101-FPN | pytorch | 1x | 6.5 | 14.4 | 40.1 | libra-faster-rcnn_r101_fpn_1x_coco.py |
| Faster R-CNN | X-101-64x4d-FPN | pytorch | 1x | 10.8 | 8.5 | 42.7 | libra-faster-rcnn_x101-64x4d_fpn_1x_coco.py |
| RetinaNet | R-50-FPN | pytorch | 1x | 4.2 | 17.7 | 37.6 | libra-retinanet_r50_fpn_1x_coco.py |
从横向对比看:在同等 Faster R-CNN R-50-FPN 条件下,Libra 版本 38.3 AP 高于标准 FPN Faster R-CNN;RetinaNet 变体 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 其他算法中,例如
sabl(SABL 采用 Bucket 回归解耦分类/回归)等目录,可与 Libra R-CNN 对照学习。
总体而言,Libra R-CNN 用三个轻量组件分别命中训练流程中的样本、特征与目标三处不平衡,MMDetection 以"继承基配置 + 最小改动"的方式将其完整落地,是理解"训练策略也是检测性能上限的一部分"这一思想的最佳入门样例之一。
【免费下载链接】mmdetectionOpenMMLab Detection Toolbox and Benchmark项目地址: https://gitcode.com/gh_mirrors/mm/mmdetection
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考