如何扩展基于Faster R-CNN的PyTorch单类目标检测程序至多类?
调整代码实现多类别目标检测
Got it, let's tweak your code to support multi-class detection with your Faster RCNN ResNet50 model. Here's the breakdown of key changes and updated code:
1. 修改数据集类以支持多类别标签
Your original RaccoonDataset hardcodes labels to torch.ones(...) since it only handles raccoons. We need to:
- Update your annotation parser (
parse_one_annot) to return both bounding boxes and their corresponding class labels - Add a class-to-ID mapping to convert string labels (like "cat", "dog") to integer IDs (note: PyTorch's detection models reserve
0for the background class, so your custom classes start at1)
Here's the updated dataset class:
import os import torch from PIL import Image class MultiClassDataset(torch.utils.data.Dataset): def __init__(self, root, data_file, transforms=None): self.root = root self.transforms = transforms self.imgs = sorted(os.listdir(os.path.join(root, "images"))) self.path_to_data_file = data_file # 自定义类别映射,根据你的数据集修改 self.class_map = { "raccoon": 1, "cat": 2, "dog": 3, # 在这里添加所有需要识别的类别,ID从1开始 } def __getitem__(self, idx): # 加载图片 img_path = os.path.join(self.root, "images", self.imgs[idx]) img = Image.open(img_path).convert("RGB") # 更新解析函数,返回(框列表, 类别名称列表) box_list, label_name_list = parse_one_annot(self.path_to_data_file, self.imgs[idx]) # 转换框为张量 boxes = torch.as_tensor(box_list, dtype=torch.float32) num_objs = len(box_list) # 将类别名称转换为整数ID labels = torch.tensor([self.class_map[name] for name in label_name_list], dtype=torch.int64) # 其余目标字段逻辑保持不变 image_id = torch.tensor([idx]) area = (boxes[:, 3] - boxes[:, 1]) * (boxes[:, 2] - boxes[:, 0]) iscrowd = torch.zeros((num_objs,), dtype=torch.int64) target = { "boxes": boxes, "labels": labels, "image_id": image_id, "area": area, "iscrowd": iscrowd } if self.transforms is not None: img, target = self.transforms(img, target) return img, target def __len__(self): return len(self.imgs)
2. 更新parse_one_annot解析函数
你需要修改parse_one_annot,让它从标注文件中提取类别标签。比如如果是VOC格式的XML标注,解析函数可以这样写(根据你的实际标注格式调整):
import xml.etree.ElementTree as ET def parse_one_annot(data_file, img_name): # 示例:VOC XML格式解析(根据你的标注路径逻辑调整) annot_path = os.path.join(os.path.dirname(data_file), "annotations", f"{os.path.splitext(img_name)[0]}.xml") tree = ET.parse(annot_path) root = tree.getroot() box_list = [] label_name_list = [] for obj in root.findall("object"): # 提取类别名称 label = obj.find("name").text label_name_list.append(label) # 提取边界框坐标 bbox = obj.find("bndbox") xmin = float(bbox.find("xmin").text) ymin = float(bbox.find("ymin").text) xmax = float(bbox.find("xmax").text) ymax = float(bbox.find("ymax").text) box_list.append([xmin, ymin, xmax, ymax]) return box_list, label_name_list
3. 调整模型初始化以匹配类别数量
加载预训练的Faster RCNN ResNet50时,需要将num_classes设置为你的自定义类别数 + 1(+1是为了包含背景类)。比如如果有3个自定义类别,num_classes=4:
import torchvision from torchvision.models.detection import fasterrcnn_resnet50_fpn # 定义总类别数(背景 + 自定义类别) num_classes = len(MultiClassDataset.class_map) + 1 # 也可以直接硬编码比如4 # 加载预训练模型并调整头部适配多类别 model = fasterrcnn_resnet50_fpn(pretrained=True) in_features = model.roi_heads.box_predictor.cls_score.in_features model.roi_heads.box_predictor = torchvision.models.detection.faster_rcnn.FastRCNNPredictor(in_features, num_classes)
关键注意事项
- 确保所有标注文件中的每个边界框都包含对应的类别标签
- 保持
class_map中的类别ID和数据集标注完全一致 - 背景类由模型自动处理,不需要在标注中添加
内容的提问来源于stack exchange,提问作者xabieririzar
相关产品推荐
相关产品推荐

