MMDetection中RTMDet大尺寸图像训练配置优化指南
📅 2026/7/27 16:37:41
👁️ 阅读次数
📝 编程学习
1. 问题背景与需求分析
在MMDetection框架中使用RTMDet算法训练自定义数据集时,遇到一个典型问题:如何调整配置文件参数以适应2448×2048的大尺寸输入图像。这在实际工业检测、医疗影像分析等场景中非常常见,因为高分辨率图像往往能保留更多细节信息。
原始配置文件默认输入尺寸通常是800×800或1333×800这类较小尺寸,直接训练大图会导致以下问题:
- 显存溢出(OOM)
- 训练速度大幅下降
- 模型收敛困难
2. 配置文件关键参数解析
2.1 数据流水线配置
在configs/_base_/datasets/coco_detection.py或类似文件中,需要修改以下关键参数:
train_pipeline = [ dict(type='LoadImageFromFile'), dict(type='LoadAnnotations', with_bbox=True), dict( type='Resize', img_scale=(2448, 2048), # 修改为目标尺寸 keep_ratio=True), # 是否保持长宽比 dict(type='RandomFlip', flip_ratio=0.5), ... ]注意:
keep_ratio=True时,实际处理会保持原图宽高比进行缩放,最终尺寸可能与设定值略有不同
2.2 模型结构配置
在configs/rtmdet/rtmdet_tiny_8xb32-300e_coco.py等模型配置文件中:
model = dict( data_preprocessor=dict( type='DetDataPreprocessor', mean=[123.675, 116.28, 103.53], # 通常不需要修改 std=[58.395, 57.12, 57.375], # 通常不需要修改 bgr_to_rgb=True, pad_size_divisor=32), # 关键参数:特征图对齐基数 backbone=dict( type='CSPNeXt', expand_ratio=0.5, deepen_factor=0.167, widen_factor=0.375, out_indices=(2, 3, 4)), neck=dict(...), bbox_head=dict( type='RTMDetHead', num_classes=80, in_channels=96, stacked_convs=2, feat_channels=96, anchor_generator=dict( type='MlvlPointGenerator', offset=0, strides=[8, 16, 32]), # 下采样率相关参数 ... ) )2.3 训练策略调整
在configs/_base_/schedules/schedule_300e.py中:
optim_wrapper = dict( type='OptimWrapper', optimizer=dict(type='AdamW', lr=0.004, weight_decay=0.05), paramwise_cfg=dict( norm_decay_mult=0, bias_decay_mult=0, bypass_duplicate=True))3. 大尺寸图像训练解决方案
3.1 显存优化策略
3.1.1 梯度累积(Gradient Accumulation)
修改configs/_base_/default_runtime.py:
train_cfg = dict( type='EpochBasedTrainLoop', max_epochs=300, val_interval=10, gradient_accumulation_steps=4) # 新增梯度累积步数3.1.2 自动混合精度(AMP)
optim_wrapper = dict( type='AmpOptimWrapper', # 修改为AMP封装器 optimizer=dict(type='AdamW', lr=0.004, weight_decay=0.05), loss_scale='dynamic')3.2 数据加载优化
3.2.1 使用多进程加载
train_dataloader = dict( batch_size=2, # 减小batch_size num_workers=8, # 增加worker数量 persistent_workers=True, sampler=dict(type='DefaultSampler', shuffle=True), batch_sampler=dict(type='AspectRatioBatchSampler'), dataset=dict(...))3.2.2 分块训练策略
对于超大图像,可考虑实现自定义Pipeline:
@TRANSFORMS.register_module() class CropLargeImage(BaseTransform): def __init__(self, crop_size=(1024, 1024), overlap=200): self.crop_size = crop_size self.overlap = overlap def transform(self, results): img = results['img'] h, w = img.shape[:2] # 实现分块逻辑 crops = [] for y in range(0, h, self.crop_size[1]-self.overlap): for x in range(0, w, self.crop_size[0]-self.overlap): crop = img[y:y+self.crop_size[1], x:x+self.crop_size[0]] crops.append(crop) # 修改results中的img和gt_bboxes results['img'] = crops results['img_shape'] = [self.crop_size]*len(crops) # 需要同步处理annotations... return results4. 参数调整经验总结
4.1 学习率调整策略
大尺寸输入时建议采用线性缩放规则(Linear Scaling Rule):
base_lr = 0.004 # 原始800x800配置 base_size = 800 * 800 new_size = 2448 * 2048 new_lr = base_lr * (new_size / base_size) # ≈0.0314.2 Anchor参数调整
对于RTMDet这类anchor-free算法,主要关注:
strides参数应与backbone下采样率匹配featmap_strides需要与neck输出特征图对应
bbox_head=dict( ... anchor_generator=dict( type='MlvlPointGenerator', strides=[8, 16, 32]), # 与backbone下采样率一致 ... )4.3 数据增强调整
大尺寸图像建议减弱空间增强强度:
train_pipeline = [ ... dict(type='RandomFlip', flip_ratio=0.3), # 降低翻转概率 dict(type='PhotoMetricDistortion', brightness_delta=32, contrast_range=(0.8, 1.2)), # 减小扰动幅度 ... ]5. 常见问题排查
5.1 显存不足(OOM)解决方案
- 减小
batch_size(最低可设为1) - 启用梯度累积(gradient_accumulation_steps)
- 使用AMP混合精度训练
- 尝试
torch.backends.cudnn.benchmark = True
5.2 训练不收敛可能原因
- 学习率未按比例放大
- 大尺寸下BatchNorm统计量不稳定
- 解决方案:使用SyncBN或GroupNorm替代
model = dict( data_preprocessor=dict(...), backbone=dict( norm_cfg=dict(type='GN', num_groups=32), # 使用GroupNorm ...), ... )5.3 验证阶段显存爆炸
可在配置文件中分离验证配置:
val_dataloader = dict( batch_size=1, # 验证时使用更小的batch_size num_workers=2, persistent_workers=True, drop_last=False, sampler=dict(type='DefaultSampler', shuffle=False), dataset=dict(...))6. 性能优化技巧
6.1 使用DALI加速数据加载
train_pipeline = [ dict(type='DALIWrapper', pipelines=[ dict(type='ImageDecoder', device='mixed'), dict(type='Resize', resize_x=2448, resize_y=2048, min_filter=types.DALIInterpType.INTERP_TRIANGULAR), ... ]), ... ]6.2 启用cudnn优化
在训练脚本开头添加:
torch.backends.cudnn.benchmark = True torch.backends.cudnn.enabled = True6.3 分布式训练配置
对于多卡训练,建议使用:
./tools/dist_train.sh \ configs/rtmdet/rtmdet_l_8xb32-300e_coco.py \ 8 # GPU数量对应修改配置文件:
optim_wrapper = dict( type='OptimWrapper', optimizer=dict(type='AdamW', lr=0.004 * 8), # 线性缩放LR ...)7. 完整配置示例
以下是适配2448×2048输入的RTMDet-L配置片段:
_base_ = './rtmdet_l_8xb32-300e_coco.py' # 数据流水线 train_pipeline = [ dict(type='LoadImageFromFile'), dict(type='LoadAnnotations', with_bbox=True), dict( type='Resize', img_scale=(2448, 2048), keep_ratio=True, interpolation='bilinear'), dict(type='RandomFlip', flip_ratio=0.3), dict(type='PhotoMetricDistortion', brightness_delta=32, contrast_range=(0.8, 1.2)), dict(type='PackDetInputs') ] # 模型调整 model = dict( data_preprocessor=dict( pad_size_divisor=64), # 增大对齐基数 backbone=dict( norm_cfg=dict(type='GN', num_groups=32)), bbox_head=dict( anchor_generator=dict( strides=[16, 32, 64]))) # 增大基础stride # 训练策略 train_dataloader = dict( batch_size=2, num_workers=8, dataset=dict(pipeline=train_pipeline)) optim_wrapper = dict( type='AmpOptimWrapper', optimizer=dict(type='AdamW', lr=0.032), clip_grad=dict(max_norm=35))
编程学习
技术分享
实战经验