如何为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
相关产品推荐
相关产品推荐

