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

如何为RetinaNet目标检测模型自动生成混淆矩阵?

为RetinaNet生成混淆矩阵的实现方案

目标检测任务的混淆矩阵和图像分类存在本质区别——图像分类是单样本单标签,而目标检测是单样本多目标,必须先完成真实框与预测框的匹配,再统计类别层面的预测/真实情况,不能直接套用分类模型的混淆矩阵代码。以下是基于你提到的inference.py和dataset.py的具体实现步骤:

核心实现步骤

1. 提取测试集的真实标注

从你的dataset.py中遍历测试集,提取每张图像中所有目标的真实类别ID,注意要保留每个独立目标的类别信息,而非整图的类别:

import torch
from dataset import CustomDataset  # 导入你的自定义数据集类

test_dataset = CustomDataset(root_dir='path/to/test', annotation_file='path/to/test_annotations.json', transform=None)

true_labels = []
for idx in range(len(test_dataset)):
    image, targets = test_dataset[idx]
    # 假设targets包含'labels'字段,存储当前图像所有目标的类别ID
    true_labels.extend(targets['labels'].numpy())

2. 获取模型的有效预测结果

修改inference.py,输出每张图像中过滤低置信度后的预测类别ID(通常置信度阈值设为0.5,可根据任务调整):

from model import retinanet  # 导入你的RetinaNet模型
import torchvision.transforms as transforms

# 加载模型与权重
model = retinanet(num_classes=你的类别总数)
model.load_state_dict(torch.load('path/to/pretrained_weights.pth'))
model.eval()

# 复用训练时的图像变换
transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])

pred_labels = []
conf_threshold = 0.5

for idx in range(len(test_dataset)):
    image, _ = test_dataset[idx]
    image_tensor = transform(image).unsqueeze(0)
    
    with torch.no_grad():
        outputs = model(image_tensor)
    
    # 过滤低置信度预测,只保留有效结果
    valid_mask = outputs['scores'] > conf_threshold
    pred_labels_batch = outputs['labels'][valid_mask]
    pred_labels.extend([label.item() for label in pred_labels_batch])

3. 匹配真实框与预测框(关键步骤)

上面的代码仅收集了所有真实和预测类别,但目标检测中需要确保每个真实框对应一个最匹配的预测框(IOU阈值通常设为0.5),避免重复统计。可以用以下方式实现匹配:

import numpy as np
from torchvision.ops import box_iou

matched_trues = []
matched_preds = []
iou_threshold = 0.5
conf_threshold = 0.5

for idx in range(len(test_dataset)):
    image, targets = test_dataset[idx]
    image_tensor = transform(image).unsqueeze(0)
    
    with torch.no_grad():
        outputs = model(image_tensor)
    
    # 过滤低置信度预测
    valid_mask = outputs['scores'] > conf_threshold
    pred_boxes = outputs['boxes'][valid_mask]
    pred_labels_batch = outputs['labels'][valid_mask]
    
    true_boxes = targets['boxes']
    true_labels_batch = targets['labels']
    
    # 计算真实框与预测框的IOU矩阵
    iou_matrix = box_iou(true_boxes, pred_boxes).numpy()
    
    # 贪心匹配真实框与预测框
    used_preds = set()
    for true_idx in range(len(true_boxes)):
        max_iou_idx = np.argmax(iou_matrix[true_idx])
        # 仅保留IOU达标且未被匹配过的预测框
        if iou_matrix[true_idx][max_iou_idx] >= iou_threshold and max_iou_idx not in used_preds:
            matched_trues.append(true_labels_batch[true_idx].item())
            matched_preds.append(pred_labels_batch[max_iou_idx].item())
            used_preds.add(max_iou_idx)
        else:
            # 无匹配预测框,视为漏检,用-1标记预测类别
            matched_trues.append(true_labels_batch[true_idx].item())
            matched_preds.append(-1)
    
    # 处理未匹配的预测框(误检)
    for pred_idx in range(len(pred_boxes)):
        if pred_idx not in used_preds:
            matched_trues.append(-1)  # 用-1标记无对应真实目标
            matched_preds.append(pred_labels_batch[pred_idx].item())

4. 生成并可视化混淆矩阵

使用sklearn工具生成混淆矩阵,可选择过滤或保留-1标记的漏检/误检情况:

from sklearn.metrics import confusion_matrix
import seaborn as sns
import matplotlib.pyplot as plt

# 过滤掉漏检/误检的标记(如需保留可跳过此步)
filtered_trues = [t for t in matched_trues if t != -1]
filtered_preds = [p for p in matched_preds if p != -1]

# 生成混淆矩阵
cm = confusion_matrix(filtered_trues, filtered_preds)
class_names = ['类别1', '类别2', ...]  # 替换为你的实际类别名称

# 可视化混淆矩阵
plt.figure(figsize=(10,8))
sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=class_names, yticklabels=class_names)
plt.xlabel('预测类别')
plt.ylabel('真实类别')
plt.title('RetinaNet 混淆矩阵')
plt.show()

注意事项

  • 置信度阈值和IOU阈值需根据任务需求调整,不同阈值会直接影响混淆矩阵的结果
  • 若数据集包含背景类,需在类别ID中对应处理,避免与漏检/误检的-1标记混淆
  • 漏检和误检可单独统计,也可在混淆矩阵中新增"无目标"类别进行标记

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 17:23:17