You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用MMDetection训练SVHN数字检测时类别数不匹配问题求助

问题描述

我想用MMDetection实现图像数字检测,采用SVHN数据集(.mat格式)作为训练样本,按照官方文档完成以下操作后:

  1. 创建了annotation.txt标注文件,格式示例:

1.png
741 350
1
43 7 19 30 5

2.png
199 83
3
99 5 14 23 2
114 8 8 23 1
121 6 12 23 0
...

  1. 编写了自定义数据集处理类:
@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']
  1. 编写了训练配置文件:
_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类场景,以下是具体修正步骤:

  1. 补全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)  # 新增该行,同步类别数
    )
)
  1. 删除冗余的数据集类型配置
    配置中同时存在dataset_type = 'COCODataset'和自定义MyDataset,会引发逻辑混淆,直接删除dataset_type = 'COCODataset'这一行即可。

  2. 验证标注标签与类别映射
    确认annotation.txt中的标签值(0-9)与MyDataset.CLASSES的索引完全对应,当前代码中直接将标注的第5个值转为int作为label,逻辑正确,无需调整。

  3. 可选:优化预训练权重加载
    加载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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.10 18:05:17