【TorchMetrics精通系列②】混淆矩阵:归一化陷阱、TP/FP推导与10分类文本热力图分析

发布时间:2026/7/20 18:04:04
【TorchMetrics精通系列②】混淆矩阵:归一化陷阱、TP/FP推导与10分类文本热力图分析 torchmetrics堪称模型评估界的“绝世秘籍”招式精妙且威力无穷。若想真正参透其中玄机、融会贯通列位看官莫急且听我细细拆解。这是 torchmetrics 系列文章的第二篇。第一篇看此处【TorchMetrics精通系列①】核心设计哲学 Accuracy 超详解聚焦在torchmetrics中的混淆矩阵。我会从概念到代码完整讲透。 一、混淆矩阵概念、样子与作用混淆矩阵Confusion Matrix是用于评估分类模型性能的表格它直观地展示了模型在每个类别上的预测结果与真实标签的对应关系。1.1 它长什么样假设一个 3 分类任务类别猫、狗、鸟模型在 100 个样本上的预测结果汇总为预测:猫预测:狗预测:鸟真实:猫2831真实:狗4252真实:鸟2530行真实标签Ground Truth列预测标签Predicted Label对角线正确分类的样本[猫→猫]28, [狗→狗]25, [鸟→鸟]30非对角线错误分类的样本能清晰看出模型把“猫”误判为“狗”3 次等等。这取决于你的目的但在torchmetrics以及绝大多数深度学习库如 sklearn、PyTorch中标准定义如下 核心口诀横真竖预横向看行 Row真实标签 (Ground Truth)竖向看列 Col预测标签 (Prediction) 具体怎么看横向看一行关注“真实”→ 看召回率 (Recall)问题“所有真实的猫模型找全了吗”怎么看盯着**“猫”的那一行**。含义这一行代表世界上所有真实的猫。对角线上的数字是找对的非对角线上的数字是漏掉的被误判成了狗或鸟。用途检查模型是否漏掉了某个类别的样本。竖向看一列关注“预测”→ 看精确率 (Precision)问题“模型预测出的猫有多少是真的”怎么看盯着**“猫”的那一列**。含义这一列代表模型信誓旦旦说是猫的所有样本。对角线上的数字是蒙对的非对角线上的数字是误报的其实是狗但被模型硬说是猫。用途检查模型是否在“指鹿为马”也就是误报率高不高。 总结想看漏没漏查全就横着看行。想看准不准查准就竖着看列。1.2 它有什么用发现类别混淆一眼看出哪些类别容易互相误判如猫 vs 狗容易混淆。计算精细指标基于混淆矩阵能推导出精确率 (Precision)、召回率 (Recall)、F1 值、特异度等更细粒度的指标。调试模型如果某两个类别间的混淆特别严重你可能需要增加这类别的训练数据或改进特征工程。对于 10 分类文本任务混淆矩阵能直接告诉你“哪些主题类别经常被模型搞混”这是单个数字如 Accuracy做不到的。 二、TorchMetrics 中的混淆矩阵TorchMetrics提供了函数式**(Functional)和模块式(Class)**两种接口来生成混淆矩阵。核心类是torchmetrics.ConfusionMatrix也可直接使用更具体的MulticlassConfusionMatrix,BinaryConfusionMatrix等子类。我们以多分类场景为主进行讲解。2.1 函数式接口签名默认值及必填/可选标注torchmetrics.functional.confusion_matrix(preds:Tensor,# 必填预测值target:Tensor,# 必填真实标签task:Literal[binary,multiclass,multilabel],# 必填任务类型num_classes:Optional[int]None,# 可选多分类时必填num_labels:Optional[int]None,# 可选多标签时必填threshold:float0.5,# 可选默认0.5二分类/多标签用normalize:Optional[Literal[true,pred,all]]None,# 可选默认None输出整数计数ignore_index:Optional[int]None,# 可选默认Nonevalidate_args:boolTrue# 可选默认True)-Tensor2.2 模块式类初始化签名默认值及必填/可选标注torchmetrics.ConfusionMatrix(task:Literal[binary,multiclass,multilabel],# 必填任务类型num_classes:Optional[int]None,# 可选多分类时必填num_labels:Optional[int]None,# 可选多标签时必填threshold:float0.5,# 可选默认0.5normalize:Optional[Literal[true,pred,all]]None,# 可选ignore_index:Optional[int]None,# 可选validate_args:boolTrue# 可选)捷径对于明确的多分类任务推荐使用torchmetrics.classification.MulticlassConfusionMatrix(num_classes10)参数更简洁不需要手动指定task。 三、参数详解参数类型必填默认值说明predsTensor✅–模型预测。可以是概率/logits浮点型或类别索引整型。多分类时若是概率形状通常为(N, C)若是类别索引形状为(N,)。targetTensor✅–真实标签。多分类时形状为(N,)的整数张量。taskLiteral[binary, multiclass, multilabel]✅–任务类型。决定混淆矩阵的维度和内部转换逻辑。num_classesOptional[int]多分类时必填None类别总数。对于 10 分类必须设为10。num_labelsOptional[int]多标签时必填None标签总数多标签任务专用。thresholdfloat可选0.5二分类或多标签时将概率转为二值预测的阈值多分类下忽略。normalizeOptional[Literal[“true”,“pred”,“all”]]可选None归一化方式•None输出原始计数值整数张量。•true按行归一化每行之和为 1即每个真实类别下预测的分布召回率视角。•pred按列归一化每列之和为 1即每个预测类别中有多少来自真实类别精确率视角。•all除以所有样本总数矩阵所有元素之和为 1。ignore_indexOptional[int]可选None指定一个类别索引计算时将其忽略该类的真实和预测都不会计入矩阵。常用于忽略填充标签。validate_argsbool可选True是否对输入参数和形状进行安全检查。 四、输入格式详解输入形式与task紧密相关针对多分类任务情况preds形状preds类型target形状target类型传入概率/logits(N, C)float32(N,)long(0 ~ C-1)传入预测类别索引(N,)long(N,)longN样本数量C类别数10如果preds是 logits未经过 softmax经过了 softmax 也可以torchmetrics内部会取argmax后再统计你无需手动转换。 五、输出结果详解形状(C, C)的矩阵其中C num_classes。数据类型normalizeNone时输出整数型torch.LongTensor原始计数。normalize为其他值时输出浮点型torch.FloatTensor。索引含义output[i, j]表示真实标签为i预测标签为j的样本数或比例。即行 真实列 预测。函数式接口直接返回该矩阵。模块式接口metric.compute()返回该矩阵metric(preds, target)会在更新状态后返回当前累积的混淆矩阵。代码示例importtorchimporttorchmetrics predstorch.tensor([0,2,1,2,0])targettorch.tensor([0,1,1,2,0])cmtorchmetrics.functional.confusion_matrix(preds,target,taskmulticlass,num_classes3)print(cm)# tensor([[2, 0, 0], # 真实02个预测为0# [0, 1, 1], # 真实11个预测为11个预测为2被误判# [0, 0, 1]]) # 真实21个预测为2若设置normalizetrue按行归一化每行之和为 1即每个真实类别下预测的分布召回率视角。cm_normtorchmetrics.functional.confusion_matrix(preds,target,taskmulticlass,num_classes3,normalizetrue)print(cm_norm)# tensor([[1.0000, 0.0000, 0.0000],# [0.0000, 0.5000, 0.5000],# [0.0000, 0.0000, 1.0000]])⚙️ 六、常用操作模块式接口6.1 基本生命周期fromtorchmetricsimportConfusionMatrix# 初始化10分类confmatConfusionMatrix(taskmulticlass,num_classes10).to(cuda)# 累积多个 batchforbatchinval_loader:preds,targetbatch confmat.update(preds,target)# 获取最终混淆矩阵cmconfmat.compute()print(cm.shape)# torch.Size([10, 10])print(cm)# 重置状态为下一轮准备confmat.reset()6.2 快捷用法仅看当前累积结果# 在训练循环内可以直接调用对象batch_cmconfmat(preds,target)# 更新并返回当前累积的混淆矩阵6.3 提取各个类别的 TP/TN/FP/FNtorchmetrics中的混淆矩阵没有直接提供提取 TP/TN/FP/FN 的高层 API但你可以基于矩阵手动计算。对于多分类通常按类别One-vs-Rest单独考虑例如对于类别iTP cm[i, i]FP cm[:, i].sum() - cm[i, i]FN cm[i, :].sum() - cm[i, i]TN cm.sum() - (TP FP FN)注意在多分类问题中TN 并不常用但公式上是有效的。6.4 可视化ConfusionMatrix对象内置了plot()方法可生成热力图需要matplotlib。该方法返回一个Figure对象你可以直接保存或记录。importtorchimporttorchmetricsimportmatplotlib.pyplotasplt# 1. 初始化10分类confmattorchmetrics.classification.MulticlassConfusionMatrix(num_classes10)# 2. 模拟累积数据for_inrange(100):predstorch.randint(0,10,(32,))targettorch.randint(0,10,(32,))confmat.update(preds,target)# 3. 绘图并保存# 【关键修改】plot() 返回的是 (fig, ax) 元组需要解包fig,axconfmat.plot()# 现在 fig 是一个 matplotlib.figure.Figure 对象可以正常保存fig.savefig(confusion_matrix.png,dpi300)# 保存图像 (dpi300 提高清晰度)plt.show()# 显示图像confmat.reset()你也可以对函数式接口的输出直接使用matplotlib自定义绘图。6.5 在 PyTorch Lightning 中记录classMyModel(pl.LightningModule):def__init__(self):super().__init__()self.confmatConfusionMatrix(taskmulticlass,num_classes10)defvalidation_step(self,batch,batch_idx):preds,targetbatch self.confmat.update(preds,target)defon_validation_epoch_end(self):cmself.confmat.compute()# 生成可视化图像并记录假设你使用 TensorBoardLoggerfigself.confmat.plot()ifself.loggerandhasattr(self.logger,experiment):self.logger.experiment.add_figure(Confusion Matrix,fig,self.current_epoch)self.confmat.reset() 总结混淆矩阵是诊断分类器错误类型的利器尤其适合 10 分类文本任务。torchmetrics中通过tasknum_classes指定支持整数计数或多种归一化输出。输入preds可以是概率矩阵或类别索引target为类别索引。输出是[num_classes, num_classes]矩阵行为真实列为预测。使用时别忘了.to(device)和reset()并且可以利用内置的plot方法直观观察。结合之前的Accuracy和Macro-F1将混淆矩阵加入评估工具链你就能既看到全局准确率又能深入到每个类别的具体表现做到“知其然也知其所以然”。