MMDetection中RTMDet大尺寸图像训练配置优化指南

MMDetection中RTMDet大尺寸图像训练配置优化指南

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 results

4. 参数调整经验总结

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.031

4.2 Anchor参数调整

对于RTMDet这类anchor-free算法,主要关注:

  1. strides参数应与backbone下采样率匹配
  2. 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)解决方案

  1. 减小batch_size(最低可设为1)
  2. 启用梯度累积(gradient_accumulation_steps)
  3. 使用AMP混合精度训练
  4. 尝试torch.backends.cudnn.benchmark = True

5.2 训练不收敛可能原因

  1. 学习率未按比例放大
  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 = True

6.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))