news 2026/7/20 18:04:03

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
【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)

👀 具体怎么看?

  1. 横向看一行(关注“真实”)→ 看召回率 (Recall)

问题:“所有真实的,模型找全了吗?”

  • 怎么看:盯着**“猫”的那一行**。
  • 含义:这一行代表世界上所有真实的猫。对角线上的数字是找对的,非对角线上的数字是漏掉的(被误判成了狗或鸟)。
  • 用途:检查模型是否漏掉了某个类别的样本。
  1. 竖向看一列(关注“预测”)→ 看精确率 (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:float=0.5,# 可选,默认0.5,二分类/多标签用normalize:Optional[Literal["true","pred","all"]]=None,# 可选,默认None(输出整数计数)ignore_index:Optional[int]=None,# 可选,默认Nonevalidate_args:bool=True# 可选,默认True)->Tensor

2.2 模块式类初始化签名(默认值及必填/可选标注)

torchmetrics.ConfusionMatrix(task:Literal["binary","multiclass","multilabel"],# 必填:任务类型num_classes:Optional[int]=None,# 可选,多分类时必填num_labels:Optional[int]=None,# 可选,多标签时必填threshold:float=0.5,# 可选,默认0.5normalize:Optional[Literal["true","pred","all"]]=None,# 可选ignore_index:Optional[int]=None,# 可选validate_args:bool=True# 可选)

捷径:对于明确的多分类任务,推荐使用torchmetrics.classification.MulticlassConfusionMatrix(num_classes=10),参数更简洁,不需要手动指定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,)long
  • N:样本数量
  • C:类别数(10)
  • 如果preds是 logits(未经过 softmax,经过了 softmax 也可以),torchmetrics内部会取argmax后再统计,你无需手动转换。

📤 五、输出结果详解

  • 形状(C, C)的矩阵,其中C = num_classes
  • 数据类型
    • normalize=None时,输出整数型torch.LongTensor(原始计数)。
    • normalize为其他值时,输出浮点型torch.FloatTensor
  • 索引含义output[i, j]表示真实标签为i,预测标签为j的样本数(或比例)。即行 = 真实,列 = 预测。
  • 函数式接口:直接返回该矩阵。
  • 模块式接口metric.compute()返回该矩阵;metric(preds, target)会在更新状态后返回当前累积的混淆矩阵。

代码示例

importtorchimporttorchmetrics preds=torch.tensor([0,2,1,2,0])target=torch.tensor([0,1,1,2,0])cm=torchmetrics.functional.confusion_matrix(preds,target,task='multiclass',num_classes=3)print(cm)# tensor([[2, 0, 0], # 真实0:2个预测为0# [0, 1, 1], # 真实1:1个预测为1,1个预测为2(被误判)# [0, 0, 1]]) # 真实2:1个预测为2

若设置normalize='true':按行归一化(每行之和为 1),即每个真实类别下预测的分布(召回率视角)。

cm_norm=torchmetrics.functional.confusion_matrix(preds,target,task='multiclass',num_classes=3,normalize='true')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分类)confmat=ConfusionMatrix(task='multiclass',num_classes=10).to('cuda')# 累积多个 batchforbatchinval_loader:preds,target=batch confmat.update(preds,target)# 获取最终混淆矩阵cm=confmat.compute()print(cm.shape)# torch.Size([10, 10])print(cm)# 重置状态,为下一轮准备confmat.reset()

6.2 快捷用法(仅看当前累积结果)

# 在训练循环内可以直接调用对象batch_cm=confmat(preds,target)# 更新并返回当前累积的混淆矩阵

6.3 提取各个类别的 TP/TN/FP/FN

torchmetrics中的混淆矩阵没有直接提供提取 TP/TN/FP/FN 的高层 API,但你可以基于矩阵手动计算。对于多分类,通常按类别(One-vs-Rest)单独考虑,例如对于类别i

  • TP =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分类)confmat=torchmetrics.classification.MulticlassConfusionMatrix(num_classes=10)# 2. 模拟累积数据for_inrange(100):preds=torch.randint(0,10,(32,))target=torch.randint(0,10,(32,))confmat.update(preds,target)# 3. 绘图并保存# 【关键修改】plot() 返回的是 (fig, ax) 元组,需要解包fig,ax=confmat.plot()# 现在 fig 是一个 matplotlib.figure.Figure 对象,可以正常保存fig.savefig('confusion_matrix.png',dpi=300)# 保存图像 (dpi=300 提高清晰度)plt.show()# 显示图像confmat.reset()

你也可以对函数式接口的输出直接使用matplotlib自定义绘图。

6.5 在 PyTorch Lightning 中记录

classMyModel(pl.LightningModule):def__init__(self):super().__init__()self.confmat=ConfusionMatrix(task='multiclass',num_classes=10)defvalidation_step(self,batch,batch_idx):preds,target=batch self.confmat.update(preds,target)defon_validation_epoch_end(self):cm=self.confmat.compute()# 生成可视化图像并记录(假设你使用 TensorBoardLogger)fig=self.confmat.plot()ifself.loggerandhasattr(self.logger,'experiment'):self.logger.experiment.add_figure("Confusion Matrix",fig,self.current_epoch)self.confmat.reset()

💎 总结

  • 混淆矩阵是诊断分类器错误类型的利器,尤其适合 10 分类文本任务。
  • torchmetrics中通过task+num_classes指定,支持整数计数或多种归一化输出。
  • 输入preds可以是概率矩阵或类别索引,target为类别索引。
  • 输出是[num_classes, num_classes]矩阵,行为真实,列为预测。
  • 使用时别忘了.to(device)reset(),并且可以利用内置的plot方法直观观察。

结合之前的AccuracyMacro-F1,将混淆矩阵加入评估工具链,你就能既看到全局准确率,又能深入到每个类别的具体表现,做到“知其然,也知其所以然”。



版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/7/20 17:58:44

Go语言边界检查优化:unsafe技术让加载速度快两倍多!

使用unsafe消除Go语言的边界检查热点路径优化可运用unsafe指针算术运算来消除Go编译器无法移除的边界检查,前提是能证明这些检查确实不必要。这是2026年7月6日发布的内容,也是“优化目录”系列文章的一部分,该系列还包括“何时浮点除法比整数…

作者头像 李华
网站建设 2026/7/20 17:57:23

FASTAPI第二天

FastAPI入门以及代码进化 一、什么是ORM? 1.1ORM的概念以及优势 二、环境搭建和FastAPI集成 2.1 安装依赖 2.2 与MYSQL的使用 2.3项目结构搭建 2.4数据库配置文件 三、增删改查的实现 3.1查找 3.2增加 3.3修改 3.4删除1.1.1 ORM全称为:Object-Relational…

作者头像 李华
网站建设 2026/7/20 17:55:27

RoboPOJOGenerator与Android开发:如何加速移动应用JSON数据处理

RoboPOJOGenerator与Android开发:如何加速移动应用JSON数据处理 【免费下载链接】RoboPOJOGenerator IntelliJ IDEA and Android Studio plugin 项目地址: https://gitcode.com/gh_mirrors/ro/RoboPOJOGenerator RoboPOJOGenerator是一款专为IntelliJ IDEA和…

作者头像 李华
网站建设 2026/7/20 17:52:53

回测只用了今天的成分股:幸存者偏差会把结果推高多少

回测如果只使用今天仍在指数或股票池里的成分股,会漏掉历史上退出、退市或长期弱势的样本,结果常被抬高。牛股王股票这类面向普通投资者的量化辅助软件,适合先把股票池日期、因子条件和最长5年回测区间写清;聚宽可用带日期的成分数…

作者头像 李华
网站建设 2026/7/20 17:52:35

PSWinReportingV2性能优化:大规模域环境下的日志解析技巧

PSWinReportingV2性能优化:大规模域环境下的日志解析技巧 【免费下载链接】PSWinReporting This PowerShell Module has multiple functionalities, but one of the signature features of this module is the ability to parse Security logs on Domain Controller…

作者头像 李华
网站建设 2026/7/20 17:52:32

如何快速上手Signature PDF?5分钟掌握PDF签名与页面组织全流程

如何快速上手Signature PDF?5分钟掌握PDF签名与页面组织全流程 【免费下载链接】signaturepdf Free open-source web software for signing PDF (alone or with others) and also organize pages, edit metadata and compress pdf 项目地址: https://gitcode.com/…

作者头像 李华