基于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
相关产品推荐
相关产品推荐

