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

基于MS COCO权重微调Mask R-CNN仅检测新增类:代码排查与修改

问题排查与修改方案

现有代码存在的核心问题

  • 配置类关键参数缺失:未指定加载COCO预训练权重的路径,且默认训练策略会更新所有层,无法避免原COCO类别权重参与训练,导致模型仍可能检测原类别。
  • 数据集加载逻辑不完整:当前load_shoe方法仅完成类别添加,未实现从数据集目录读取图片、解析标注(如掩码、边界框)等核心数据加载逻辑,模型无法获取有效训练样本。
  • 类别映射未处理:直接使用新类别ID(1对应shoe)与COCO预训练模型的类别ID体系不匹配,会导致权重加载混乱。

修改后的完整代码方案

1. 配置类修改

class ShoeConfig(Config):
    NAME = "shoe"

    IMAGES_PER_GPU = 2

    # 类别数(背景+新增的shoe类)
    NUM_CLASSES = 1 + 1  # Background + shoe

    STEPS_PER_EPOCH = 11797

    DETECTION_MIN_CONFIDENCE = 0.9

    # 指定加载COCO预训练权重文件路径
    WEIGHTS = "mask_rcnn_coco.h5"

    # 仅训练头部层(分类、掩码、边界框回归分支),冻结主干网络及原COCO类别相关权重
    TRAINABLE_LAYERS = "heads"

2. 数据集加载代码补全(适配COCO格式标注)

import json
import numpy as np
import cv2

class ShoeDataset(utils.Dataset):

    def load_shoe(self, dataset_dir, subset):
        # 添加目标类别
        self.add_class("shoe", 1, "shoe")

        # 确认子集类型
        assert subset in ["train", "val"]
        dataset_dir = os.path.join(dataset_dir, subset)

        # 读取COCO格式标注文件(假设标注文件名为instances_shoe.json)
        annotation_path = os.path.join(dataset_dir, "instances_shoe.json")
        with open(annotation_path, 'r') as f:
            coco_json = json.load(f)

        # 建立图片ID到标注的映射
        img_to_anns = {}
        for ann in coco_json['annotations']:
            img_id = ann['image_id']
            img_to_anns.setdefault(img_id, []).append(ann)

        # 遍历图片信息,添加到数据集
        for img in coco_json['images']:
            img_id = img['id']
            img_path = os.path.join(dataset_dir, img['file_name'])
            self.add_image(
                "shoe",
                image_id=img_id,
                path=img_path,
                width=img['width'],
                height=img['height'],
                annotations=img_to_anns.get(img_id, [])
            )

    def load_mask(self, image_id):
        # 生成图片对应的掩码矩阵
        image_info = self.image_info[image_id]
        annotations = image_info['annotations']
        count = len(annotations)
        mask = np.zeros([image_info['height'], image_info['width'], count], dtype=np.uint8)
        class_ids = []

        for i, ann in enumerate(annotations):
            # 解析多边形标注并填充掩码
            for seg in ann['segmentation']:
                seg_coords = np.array(seg).reshape((-1, 2)).astype(np.int32)
                cv2.fillPoly(mask[:, :, i], [seg_coords], 1)
            # 指定当前实例的类别ID
            class_ids.append(self.class_names.index("shoe"))
        
        return mask, np.array(class_ids, dtype=np.int32)

    def image_reference(self, image_id):
        # 返回图片路径作为参考
        return self.image_info[image_id]['path']

额外注意事项

  • 若你的数据集标注不是COCO格式,需对应修改load_shoe和load_mask方法中的标注解析逻辑。
  • 加载预训练权重时,因原COCO模型有80类、当前仅1类,会出现权重不匹配警告,可直接忽略——因为我们已冻结非头部层,这些不匹配的权重不会被更新。
  • 训练完成后,模型仅会输出背景与shoe类的检测结果,不会识别原COCO的其他类别。

内容的提问来源于stack exchange,提问作者Eave_Z

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 18:50:25