使用MMDetection训练SVHN数字检测时类别数不匹配问题求助
问题描述
我想用MMDetection实现图像数字检测,采用SVHN数据集(.mat格式)作为训练样本,按照官方文档完成以下操作后:
- 创建了
annotation.txt标注文件,格式示例:
1.png
741 350
1
43 7 19 30 52.png
199 83
3
99 5 14 23 2
114 8 8 23 1
121 6 12 23 0
...
- 编写了自定义数据集处理类:
@DATASETS.register_module() class MyDataset(CustomDataset): CLASSES = ('0', '1', '2', '3', '4', '5', '6', '7', '8', '9') def load_annotations(self, ann_file): ann_list = mmcv.list_from_file(ann_file) data_infos = [] for i, ann_line in enumerate(ann_list): if ann_line != '#': continue img_shape = ann_list[i + 2].split(' ') width = int(img_shape[0]) height = int(img_shape[1]) bbox_number = int(ann_list[i + 3]) bboxes = [] labels = [] for anns in ann_list[i + 4:i + 4 + bbox_number]: anns = anns.strip().split(' ') bboxes.append([float(ann) for ann in anns[:4]]) labels.append(int(anns[4])) data_infos.append( dict( filename=ann_list[i + 1], width=width, height=height, ann=dict( bboxes=np.array(bboxes).astype(np.float32), labels=np.array(labels).astype(np.int64)) )) return data_infos def get_ann_info(self, idx): return self.data_infos[idx]['ann']
- 编写了训练配置文件:
_base_ = '../mask_rcnn/mask_rcnn_r50_caffe_fpn_mstrain-poly_1x_coco.py' model = dict( roi_head=dict( bbox_head=dict(num_classes=10))) img_norm_cfg = dict( mean=[123.675, 116.28, 103.53], std=[58.395, 57.12, 57.375], to_rgb=True) train_pipeline = [ dict(type='LoadImageFromFile'), dict(type='LoadAnnotations', with_bbox=True, with_mask=False), dict( type='Resize', img_scale=[(1333, 640), (1333, 672), (1333, 704), (1333, 736), (1333, 768), (1333, 800)], multiscale_mode="value", keep_ratio=True), dict(type='RandomFlip', flip_ratio=0.5), dict(type='Normalize', **img_norm_cfg), dict(type='Pad', size_divisor=32), dict(type='DefaultFormatBundle'), dict(type='Collect', keys=['img', 'gt_bboxes', 'gt_labels']), ] test_pipeline = [ dict(type='LoadImageFromFile'), dict( type='MultiScaleFlipAug', img_scale=(1333, 800), flip=False, transforms=[ dict(type='Resize', keep_ratio=True), dict(type='RandomFlip'), dict(type='Normalize', **img_norm_cfg), dict(type='Pad', size_divisor=32), dict(type='ImageToTensor', keys=['img']), dict(type='Collect', keys=['img']), ]) ] dataset_type = 'COCODataset' classes = ('0', '1', '2', '3', '4', '5', '6', '7', '8', '9',) data = dict( train=dict( img_prefix='../dataset/train/', type='MyDataset', ann_file='../annotation_train.txt', pipeline=train_pipeline), val=dict( img_prefix='../dataset/test/', type='MyDataset', ann_file='../annotation_test.txt', pipeline=test_pipeline), test=dict( img_prefix='../dataset/test/', type='MyDataset', ann_file='../annotation_test.txt', pipeline=test_pipeline)) load_from = 'checkpoints/mask_rcnn_r50_caffe_fpn_mstrain-poly_2x_coco_bbox_mAP-0.403__segm_mAP-0.365_20200504_231822-a75c98ce.pth'
运行时出现错误,提示类别数应为80而非10,求解决方法。
解决方案
错误核心是继承的COCO基础配置(默认80类)未完全适配自定义10类场景,以下是具体修正步骤:
- 补全Mask R-CNN的mask_head类别配置
Mask R-CNN包含bbox和mask两个分支,你仅修改了bbox_head的num_classes,但mask_head仍保持默认80类,导致不匹配。修改模型配置部分:
model = dict( roi_head=dict( bbox_head=dict(num_classes=10), mask_head=dict(num_classes=10) # 新增该行,同步类别数 ) )
删除冗余的数据集类型配置
配置中同时存在dataset_type = 'COCODataset'和自定义MyDataset,会引发逻辑混淆,直接删除dataset_type = 'COCODataset'这一行即可。验证标注标签与类别映射
确认annotation.txt中的标签值(0-9)与MyDataset.CLASSES的索引完全对应,当前代码中直接将标注的第5个值转为int作为label,逻辑正确,无需调整。可选:优化预训练权重加载
加载COCO预训练权重时,MMDetection会自动用随机初始化替换原80类的头部权重。若后续训练效果不佳,可指定仅加载backbone权重:
init_cfg = dict( type='Pretrained', checkpoint='checkpoints/mask_rcnn_r50_caffe_fpn_mstrain-poly_2x_coco_bbox_mAP-0.403__segm_mAP-0.365_20200504_231822-a75c98ce.pth', prefix='backbone.' )
内容的提问来源于stack exchange,提问作者Eugeny Bessonn
相关产品推荐
相关产品推荐

