尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

目标检测模型评估利器:多类混淆矩阵原理、实现与优化指南

目标检测模型评估利器:多类混淆矩阵原理、实现与优化指南 1. 项目概述为什么我们需要“多类混淆矩阵”在目标检测这个领域里摸爬滚打了这么多年我见过太多人把模型训练出来一看mAP平均精度挺高就兴冲冲地准备上线了。结果一到真实场景模型对“猫”和“狗”的识别总是打架或者把“卡车”认成“公交车”搞得下游应用一团糟。这时候一个简单的“准确率”或者“召回率”数字就像一份只告诉你“总分”却不给“各科成绩”的成绩单根本没法帮你定位问题到底出在哪里。这就是“多类混淆矩阵”这个工具的价值所在——它不是给你一个最终答案而是给你一张清晰的“诊断报告”让你能看清模型在每个类别上的“爱恨情仇”。简单来说多类混淆矩阵是评估多类别分类模型包括目标检测中的分类部分性能的核心工具。它不再满足于“对”或“错”的二元判断而是将预测结果和真实标签进行交叉比对形成一个N x N的表格N是类别数。表格的每一行代表一个真实的类别每一列代表模型预测的类别。对角线上的数字就是模型预测正确的样本数而其他格子里的数字则清晰地揭示了模型犯下的每一种“错误”哪些类别容易被混淆谁经常被误认为谁这些信息是任何单一指标都无法提供的。对于目标检测任务情况会稍微复杂一点。因为目标检测不仅要判断“是什么”分类还要定位“在哪里”定位。所以我们通常会在计算混淆矩阵之前先基于IoU交并比阈值来判定一个检测框是否匹配上了真实框。只有匹配上的检测我们才会去分析它的分类对错从而填入混淆矩阵。这就像先确认你找到了一个“目标”再去判断你叫对了它的“名字”。理解了这一点你就能明白为什么一个设计良好的多类混淆矩阵是优化目标检测模型、解决类别不平衡、提升模型鲁棒性的必备利器。2. 核心概念与目标检测评估背景2.1 从二分类到多分类混淆矩阵的演进要理解多类混淆矩阵最好从它的源头——二分类混淆矩阵说起。假设我们只有一个“猫”类。那么矩阵就是一个2x2的表格真正例TP真实是猫预测也是猫。假正例FP真实不是猫是背景或其他预测是猫。假反例FN真实是猫预测不是猫或没检测到。真反例TN真实不是猫预测也不是猫。从这个简单的表格里我们能衍生出准确率、精确率、召回率、F1分数等一系列指标。但是当类别变成“猫”、“狗”、“车”等多个时TN真反例这个概念就变得模糊了。对于“猫”类来说一个被正确预测为“狗”的样本算不算“真反例”呢从“猫”的角度看它确实不是猫预测也不是猫似乎算TN。但这会带来统计上的重复和混乱。因此多类混淆矩阵摒弃了全局的TN概念转而聚焦于每个类别的TP、FP和FN。计算方式通常有两种主流思路“一对多”模式对于类别C将所有其他类别视为“非C”。这样对于C类我们可以计算出一个二元的TP_C, FP_C, FN_C。将所有类别的这些数据以矩阵形式呈现就是多类混淆矩阵。这是最常用、最直观的理解方式。直接统计模式直接统计所有样本对于每个样本根据其真实标签i和预测标签j在矩阵的(i, j)位置加1。最终矩阵的每一行之和等于该类别的真实样本数每一列之和等于被预测为该类别的样本数。2.2 目标检测的特殊性分类与定位的耦合目标检测的评估比图像分类复杂一个维度因为它耦合了“定位”和“分类”两个任务。你不能因为分类对了就认为检测成功反之亦然。这就引入了两个关键概念交并比IoU这是衡量预测框B_pred和真实框B_gt重叠程度的指标。计算公式为IoU Area(B_pred ∩ B_gt) / Area(B_pred ∪ B_gt)。它的值在0到1之间越接近1说明定位越准。置信度Confidence Score模型对每个预测框所包含物体属于某类别的确信程度也是一个0到1之间的分数。在评估时我们通常遵循以下流程对每一张图片根据置信度对所有预测框进行排序例如从高到低。对于每个真实框在所有预测框中寻找与其IoU大于某个阈值常用0.5即PASCAL VOC标准更严格的如COCO使用0.5:0.05:0.95的多阈值平均且置信度最高的预测框进行匹配。匹配成功的预测框如果其预测类别与真实类别一致则计为该类的一个TP如果类别不一致则计为双重错误对于真实类别这是一个FN漏检对于预测类别这是一个FP误检。未能与任何真实框匹配的预测框计为FP虚警。未能与任何预测框匹配的真实框计为FN漏检。注意这里有一个关键细节即“非极大值抑制NMS”。NMS是在模型推理后、评估前进行的后处理步骤用于消除同一个物体上的多个重叠框。评估用的预测框必须是经过NMS处理后的结果否则会因大量重复框导致FP激增混淆矩阵失去意义。2.3 多类混淆矩阵在目标检测中的构建逻辑结合上述背景构建用于目标检测的多类混淆矩阵其核心逻辑如下确定评估粒度我们是在“图片级别”还是“实例框级别”统计目标检测中都是在“实例级别”。即每一个真实标注框和每一个预测框都是一个统计样本。基于IoU匹配只有与真实框成功匹配IoU 阈值的预测框才会参与分类对错的统计并填入混淆矩阵。填充矩阵假设真实类别为i预测类别为j。如果i j则在矩阵的(i, i)位置对角线计数加1。这是一个TP。如果i ! j则在矩阵的(i, j)位置计数加1。这表示一个分类错误对真实类别i而言这是一个FN对预测类别j而言这是一个FP。处理未匹配项所有未匹配的真实框每个都为其真实类别i增加一个FN。这体现在矩阵上是该行类别i的“总和”与统计到的TP分类错误数之间的差额。所有未匹配的预测框每个都为其预测类别j增加一个FP。这体现在矩阵上是该列类别j的“总和”与统计到的TP分类错误数之间的差额。因此一个完整的目标检测多类混淆矩阵其行和真实数量 TP_i FN_i含分类错误和未匹配的漏检列和预测数量 TP_j FP_j含分类错误和未匹配的虚警。通过这个矩阵我们可以一目了然地看到对角线各类别的检测能力TP。非对角线的密集区域类别之间的混淆情况。行合计与对角线的差距该类别的漏检情况。列合计与对角线的差距该类别的虚警情况。3. 代码实现与可视化实战理解了原理我们动手实现一个。这里以PyTorch环境下基于YOLOv8模型输出和COCO格式数据集为例进行说明。我们会用到torchmetrics库它提供了高度封装且高效的ConfusionMatrix实现。3.1 环境准备与数据假设首先确保你的环境已安装必要库。pip install torch torchvision torchmetrics seaborn matplotlib假设我们有一个简单的验证数据加载流程能返回预测和标签。预测格式通常为(batch_size, N_pred, 6)其中最后一维的6表示[x1, y1, x2, y2, confidence, class_id]。标签格式为(batch_size, N_gt, 5)其中最后一维的5表示[class_id, x1, y1, x2, y2]COCO格式通常class_id在前。为了简化我们定义一个函数来模拟一个批次的评估过程。3.2 核心计算步骤拆解我们不直接调用高级API先拆解每一步看看混淆矩阵是如何从一堆框里算出来的。import torch import numpy as np def calculate_detection_confusion_matrix_for_batch(preds, targets, iou_threshold0.5, num_classes80): 计算一个批次的目标检测混淆矩阵。 preds: List[Tensor]每个元素是 (N_pred, 6) [x1, y1, x2, y2, conf, cls] targets: List[Tensor]每个元素是 (N_gt, 5) [cls, x1, y1, x2, y2] iou_threshold: IoU匹配阈值 num_classes: 总类别数 返回: 本批次累积的混淆矩阵 (num_classes, num_classes) # 初始化混淆矩阵 cm torch.zeros((num_classes, num_classes), dtypetorch.int64) # 逐张图片处理 for pred, target in zip(preds, targets): if len(target) 0: # 图片中没有真实目标 # 所有预测都是FP按预测类别计入混淆矩阵的“列” if len(pred) 0: fp_classes pred[:, 5].long() # 预测类别ID for cls in fp_classes: # 对于FP真实类别是“背景”但我们矩阵不包含背景类。 # 一种处理方式是将其视为真实类别是“背景”导致的FN但背景不在矩阵内。 # 更标准的做法是在计算各类别FP时单独累计。这里为简化我们在矩阵上不直接体现。 # 实际上我们需要一个额外的数组来记录每个类别的FP未匹配预测。 pass # 具体处理见下方完整流程 continue if len(pred) 0: # 图片中没有预测 # 所有真实目标都是FN按真实类别计入混淆矩阵的“行” fn_classes target[:, 0].long() for cls in fn_classes: # FN意味着真实类别是cls但预测成了“背景”或其他未匹配。 # 在矩阵中这体现为该行(cls)的统计数少于真实数量。我们稍后统一处理。 pass continue # 1. 计算所有预测框和所有真实框之间的IoU矩阵 (M_pred, M_gt) iou_matrix box_iou(pred[:, :4], target[:, 1:5]) # 假设box_iou函数已实现 # 2. 根据IoU阈值和置信度进行匹配贪心算法常见于评估 # 通常按置信度降序处理预测框 conf, conf_idx pred[:, 4].sort(descendingTrue) pred pred[conf_idx] iou_matrix iou_matrix[conf_idx] matched_gt set() for pred_idx, single_pred in enumerate(pred): # 找到与该预测框IoU最大的真实框 ious iou_matrix[pred_idx] if len(ious) 0: continue max_iou, gt_idx ious.max(dim0) if max_iou iou_threshold and gt_idx not in matched_gt: # 匹配成功 matched_gt.add(gt_idx) gt_cls target[gt_idx, 0].long() pred_cls single_pred[5].long() cm[gt_cls, pred_cls] 1 # 填充混淆矩阵 else: # 未匹配的预测框 - FP预测类别为 pred_cls pred_cls single_pred[5].long() # 记录到FP计数这里我们先省略最后统一处理 pass # 3. 处理未匹配的真实框 - FN all_gt_indices set(range(len(target))) unmatched_gt_indices all_gt_indices - matched_gt for gt_idx in unmatched_gt_indices: gt_cls target[gt_idx, 0].long() # FN意味着真实类别gt_cls被预测为背景或其他在矩阵中体现为行合计不足。 # 我们需要知道总真实数所以这里先记录最后通过比较行和与真实数量来计算FN。 pass # 注意上述简化流程没有完整维护FP和FN的矩阵映射实际完整的计算需要更复杂的簿记。 # 下面展示使用 torchmetrics 的完整方案。 return cm def box_iou(boxes1, boxes2): 计算两个框集合之间的IoU矩阵。 # 简化实现实际需考虑广播和面积计算 # 此处为示意省略详细实现 pass实操心得自己从零实现匹配和矩阵填充是一个很好的学习过程能让你深刻理解评估的每个细节。但在生产或研究中强烈建议使用成熟的库如torchmetrics.detection模块下的MeanAveragePrecision它内部已经集成了混淆矩阵的计算并且经过了高度优化避免了自己实现可能带来的边界条件错误和性能瓶颈。3.3 使用TorchMetrics进行高效计算torchmetrics库的ConfusionMatrix任务是为分类设计的但我们可以利用其“多分类”模式并确保输入是经过目标检测匹配后的“分类结果”。from torchmetrics.detection import MeanAveragePrecision from torchmetrics.classification import MulticlassConfusionMatrix import torch # 假设我们有模型在整个验证集上的预测结果和真实标签 # all_preds: List[Dict]每个Dict {boxes: Tensor(N,4), scores: Tensor(N,), labels: Tensor(N,)} # all_targets: List[Dict]每个Dict {boxes: Tensor(M,4), labels: Tensor(M,)} # 1. 首先用MeanAveragePrecision计算mAP同时它内部完成了匹配 map_metric MeanAveragePrecision(iou_thresholds[0.5], box_formatxyxy) # 假设我们遍历数据集逐批更新metric # for pred_batch, target_batch in val_loader: # map_metric.update(pred_batch, target_batch) # 2. 但我们需要混淆矩阵。更直接的方法是自己执行匹配然后收集分类结果。 def collect_cls_pairs_for_confusion_matrix(all_preds, all_targets, iou_thresh0.5): 收集所有匹配成功的真实类别预测类别对用于构建混淆矩阵。 同时收集未匹配的真实类别FN和未匹配的预测类别FP信息。 all_true_labels [] all_pred_labels [] # 用于记录每个类别的FP和FN可选用于分析 fp_by_class torch.zeros(num_classes, dtypetorch.long) fn_by_class torch.zeros(num_classes, dtypetorch.long) for pred_dict, target_dict in zip(all_preds, all_targets): pred_boxes pred_dict[boxes] pred_scores pred_dict[scores] pred_labels pred_dict[labels] gt_boxes target_dict[boxes] gt_labels target_dict[labels] if len(gt_boxes) 0: # 没有真实框所有预测都是FP fp_by_class.index_add_(0, pred_labels, torch.ones_like(pred_labels, dtypetorch.long)) continue if len(pred_boxes) 0: # 没有预测框所有真实都是FN fn_by_class.index_add_(0, gt_labels, torch.ones_like(gt_labels, dtypetorch.long)) continue # 计算IoU矩阵 (M_pred, M_gt) iou_matrix box_iou(pred_boxes, gt_boxes) # 需要实际的box_iou实现 # 按置信度排序预测 if len(pred_scores) 0: conf_sorted_idx torch.argsort(pred_scores, descendingTrue) pred_boxes pred_boxes[conf_sorted_idx] pred_labels pred_labels[conf_sorted_idx] iou_matrix iou_matrix[conf_sorted_idx] matched_gt_indices set() for pred_idx, (pred_box, pred_label) in enumerate(zip(pred_boxes, pred_labels)): ious iou_matrix[pred_idx] if ious.numel() 0: max_iou, gt_idx ious.max(dim0) if max_iou iou_thresh and gt_idx not in matched_gt_indices: # 匹配成功 matched_gt_indices.add(gt_idx) gt_label gt_labels[gt_idx] all_true_labels.append(gt_label.item()) all_pred_labels.append(pred_label.item()) else: # 未匹配预测 - FP fp_by_class[pred_label] 1 else: fp_by_class[pred_label] 1 # 处理未匹配的真实框 - FN for gt_idx, gt_label in enumerate(gt_labels): if gt_idx not in matched_gt_indices: fn_by_class[gt_label] 1 return all_true_labels, all_pred_labels, fp_by_class, fn_by_class # 3. 使用收集到的标签对创建混淆矩阵 num_classes 80 # 例如COCO有80类 true_labels, pred_labels, fp_by_cls, fn_by_cls collect_cls_pairs_for_confusion_matrix(all_preds, all_targets) # 将列表转换为Tensor true_tensor torch.tensor(true_labels, dtypetorch.long) pred_tensor torch.tensor(pred_labels, dtypetorch.long) # 初始化并计算混淆矩阵 confmat MulticlassConfusionMatrix(num_classesnum_classes, normalizenone) cm_tensor confmat(pred_tensor, true_tensor) # 注意参数顺序preds, target print(f混淆矩阵形状: {cm_tensor.shape}) print(cm_tensor)3.4 结果可视化让问题一目了然一个数字矩阵不直观我们需要可视化。通常使用热力图并用颜色深浅表示数量多少。import matplotlib.pyplot as plt import seaborn as sns import numpy as np def plot_confusion_matrix(cm_array, class_names, figsize(20, 16), normalizeTrue): cm_array: numpy array, shape (num_classes, num_classes) class_names: list of string, 类别名称列表 normalize: 是否按行真实类别归一化显示的是召回率比例 if normalize: # 按行归一化即每个真实类别的预测分布 cm_normalized cm_array.astype(float) / (cm_array.sum(axis1)[:, np.newaxis] 1e-6) # 将对角线TP之外的部分高亮显示错误 np.fill_diagonal(cm_normalized, 0) # 可选将对角线置零以更清晰查看错误 data_to_plot cm_normalized fmt .2f title_suffix (Normalized by Row/Recall) else: data_to_plot cm_array fmt d title_suffix (Raw Counts) plt.figure(figsizefigsize) # 使用seaborn绘制热力图 sns.heatmap(data_to_plot, annotFalse, fmtfmt, cmapBlues, # 或者Reds查看错误 xticklabelsclass_names, yticklabelsclass_names, cbar_kws{label: 比例 if normalize else 数量}) plt.title(f目标检测多类混淆矩阵{title_suffix}) plt.xlabel(预测类别) plt.ylabel(真实类别) plt.xticks(rotation90, haright) plt.yticks(rotation0) plt.tight_layout() plt.show() # 假设我们有类别名称列表 coco_names # cm_numpy cm_tensor.numpy() # plot_confusion_matrix(cm_numpy, coco_names, normalizeTrue) # 查看原始计数 # plot_confusion_matrix(cm_numpy, coco_names, normalizeFalse)注意事项可视化时归一化矩阵按行非常有用它直接显示了对于每个真实类别模型将其预测为各个类别的概率分布。颜色越深非对角线的地方就是混淆最严重的地方。例如如果“猫”所在的行在“狗”的列上有一个深色块就意味着很多猫被误认成了狗。4. 从混淆矩阵中挖掘洞察与模型优化得到了混淆矩阵工作才完成了一半。更重要的是如何解读它并指导模型优化。4.1 关键指标计算与解读从混淆矩阵中我们可以为每个类别i计算出更细致的指标各类别精确率Precision per Class:P_i TP_i / (TP_i FP_i) cm[i,i] / (cm[:, i].sum())解释在所有被预测为类别i的框中有多少是真的i。低精确率说明该类虚警多。各类别召回率Recall per Class:R_i TP_i / (TP_i FN_i) cm[i,i] / (cm[i, :].sum())解释在所有真实的类别i的框中有多少被成功检测并正确分类。低召回率说明该类漏检多。各类别F1分数:F1_i 2 * (P_i * R_i) / (P_i R_i)我们可以编写一个函数来自动计算这些值def analyze_confusion_matrix(cm_tensor, class_names): cm_tensor: torch.Tensor, (C, C) class_names: list of length C cm cm_tensor.numpy() num_classes cm.shape[0] analysis [] for i in range(num_classes): tp cm[i, i] fp cm[:, i].sum() - tp # 预测为i但真实不是i的总数含背景未匹配的FP需额外加 fn cm[i, :].sum() - tp # 真实是i但预测不是i的总数含未匹配的FN需额外加 # 注意这里计算的是基于矩阵内匹配对的FP和FN未包含完全未匹配的FP/FN。 # 更严谨的做法需要传入之前单独统计的 fp_by_class 和 fn_by_class。 precision tp / (tp fp 1e-9) recall tp / (tp fn 1e-9) f1 2 * precision * recall / (precision recall 1e-9) analysis.append({ class: class_names[i], TP: int(tp), FP: int(fp), FN: int(fn), Precision: round(precision, 4), Recall: round(recall, 4), F1: round(f1, 4) }) # 转换为DataFrame便于查看 import pandas as pd df pd.DataFrame(analysis) df df.sort_values(F1, ascendingTrue) # 按F1排序找出最差的类别 print(df.head(20)) # 查看表现最差的20个类别 return df4.2 典型问题模式与解决方案通过观察混淆矩阵和上述分析表格你可以识别出几种常见问题模式特定类别召回率极低行和很大但对角线值很小问题模型根本检测不到这个类别或者检测到了但分类全错。排查查看该类别训练样本数量是否严重不足类别不平衡。检查标注质量该类别物体是否特别小、遮挡严重或形状特殊。优化数据层面收集更多该类别数据使用过采样如复制、数据增强针对小目标的随机裁剪、缩放、拼接Mosaic。模型层面调整锚框Anchor尺寸使其更匹配小目标使用专门针对小目标改进的检测头如添加高分辨率特征图分支尝试FPN、PANet等加强特征金字塔网络。损失函数使用Focal Loss缓解类别不平衡让模型更关注难例。特定类别精确率极低列和很大但对角线值很小问题模型经常把其他东西误认为这个类别虚警高。排查观察该列预测类别下哪些真实类别的贡献最大混淆矩阵中该列的非对角线高亮行。例如“狗”的列在“猫”的行上值很高说明猫狗混淆。优化数据与特征检查混淆类别之间的训练数据是否外观相似如“狼”和“哈士奇”。增加能区分这两类特征的数据。后处理提高该类别在NMS中的置信度阈值或IoU阈值过滤掉一些低质量的预测。模型结构考虑在分类头引入更强大的特征提取器或者使用解耦头Decoupled Head分别优化分类和回归任务。对称性混淆两个类别互相误认问题例如“苹果”和“橘子”互相认错。排查这是典型的特征相似性问题。查看这两个类别的训练样本在特征空间是否距离很近。优化数据增强使用CutMix、MixUp等混合类别的增强方式让模型学习更鲁棒的边界。度量学习在损失函数中加入对比损失Contrastive Loss或三元组损失Triplet Loss拉大同类别特征距离推远不同类别特征距离。集成训练一个专门的二分类器来区分这两个易混淆类别作为后处理。某一类被预测为多种其他类一行中多个非对角线格有值问题例如“椅子”被误认为“沙发”、“凳子”、“桌子”等多种家具。排查这可能是该类别的定义模糊或类内差异大。“椅子”本身形态多样。优化考虑是否需要进行类别合并或重新定义。或者引入更细粒度的子类别。4.3 利用混淆矩阵指导数据标注与清洗混淆矩阵不仅是模型诊断工具也是数据质量的“照妖镜”。发现标注错误如果某个样本的真实类别是A但模型以极高置信度预测为B并且A和B在混淆矩阵中显示易混淆那么很可能是这个样本的标注错了。你可以导出这些高置信度错误样本进行人工复核。指导主动学习在模型不确定预测概率分布平缓或高置信度错误预测为B但真实是A且B列A行值高的区域进行采样对这些样本进行标注投入训练性价比最高。评估数据增强策略在应用了新的数据增强如针对小目标的随机粘贴后重新计算混淆矩阵观察之前召回率低的类别是否有所改善同时要警惕是否引入了新的混淆精确率下降。5. 高级话题与工程实践要点5.1 处理背景类与“未知”类别在目标检测中背景是一个特殊的“类别”。但标准的混淆矩阵通常不包含背景类因为背景是无穷多的。未匹配的预测框FP本质就是模型将背景预测为前景物体。为了更全面分析可以在计算各类别精确率时分母TP_i FP_i中的FP_i必须包含这些未匹配的预测框。可以单独维护一个FP_background统计但更常见的做法是将其分摊到各个预测类别j的FP_j中正如我们代码里fp_by_class所做的。对于开放世界目标检测模型可能会遇到训练集中未出现的“未知”类别。一个健壮的系统需要能识别出“我不知道这是什么”。这通常通过设定一个置信度阈值低于阈值的预测被归为“背景”或“未知”。评估时可以将“未知”作为一个单独的类别加入混淆矩阵分析模型对已知类别的拒绝能力和对未知类别的识别能力。5.2 在多阶段检测器如Faster R-CNN中的应用对于两阶段检测器混淆矩阵的分析可以更细致地定位问题出在哪一阶段第一阶段RPN的问题如果某个类别的召回率很低但被提议出来的区域Region Proposal其实包含了该物体那么问题可能出在RPN的召回上。可以分析RPN对该类别的提议召回率。第二阶段RoI Head的问题如果提议很多但分类错误混淆矩阵非对角线值高那么问题主要在RoI Head的分类器。可以单独提取RPN提议的特征查看分类器在这些特征上的混淆矩阵。5.3 与mAP指标的关联与区别mAP平均精度是目标检测的权威综合指标它综合了不同IoU阈值和召回率下的精确率。混淆矩阵是mAP计算过程中的一个“中间产物”或“更细致的展现”。mAP给你一个总分告诉你模型整体好坏。混淆矩阵给你一份详细的错题集告诉你哪里错了怎么错的。在模型迭代中应该两者结合看用mAP跟踪整体进度用混淆矩阵指导具体优化方向。当mAP卡住时混淆矩阵能提供突破的线索。5.4 分布式评估与大规模数据集的处理当数据集非常大如COCO时逐张图片计算并累积匹配关系可能消耗大量内存。工程上需要注意流式处理逐批处理数据只累积最终的计数矩阵或匹配对列表而不是保存所有中间匹配结果。使用高效IoU计算使用向量化操作和GPU加速的IoU计算库如torchvision.ops.box_iou。利用成熟框架的评估器像MMDetection、Detectron2等框架其评估模块都经过了高度优化应优先使用。它们通常提供了计算混淆矩阵的接口或可以轻松扩展。最后记住一点混淆矩阵是一个强大的诊断工具但它反映的是模型在当前测试集上的表现。要确保测试集本身没有偏差并且覆盖了实际应用场景的多样性。定期根据混淆矩阵的发现去审视和丰富你的数据集才能让模型越练越强。在我自己的项目中每次模型迭代后的第一件事就是生成并仔细研读这份“错题集”它往往比任何自动调参算法都更能直击要害。
返回列表