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

如何扩展基于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 0 for the background class, so your custom classes start at 1)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.09 07:22:29