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

torchmetrics 堪称模型评估界的"绝世秘籍",招式精妙且威力无穷。若想真正参透其中玄机、融会贯通,列位看官莫急,且听我细细拆解。

这是 torchmetrics 系列文章的第二篇。

第一篇看此处:【TorchMetrics精通系列①】核心设计哲学 + Accuracy 超详解

聚焦在 torchmetrics 中的混淆矩阵。我会从概念到代码,完整讲透。

📊 一、混淆矩阵:概念、样子与作用

混淆矩阵(Confusion Matrix)是用于评估分类模型性能的表格,它直观地展示了模型在每个类别上的预测结果与真实标签的对应关系。

1.1 它长什么样?

假设一个 3 分类任务(类别:猫、狗、鸟),模型在 100 个样本上的预测结果汇总为:

预测:猫 预测:狗 预测:鸟
真实:猫 28 3 1
真实:狗 4 25 2
真实:鸟 2 5 30
  • :真实标签(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 函数式接口签名(默认值及必填/可选标注)

python 复制代码
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,         # 可选,默认None
    validate_args: bool = True                  # 可选,默认True
) -> Tensor

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

python 复制代码
torchmetrics.ConfusionMatrix(
    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,  # 可选
    ignore_index: Optional[int] = None,         # 可选
    validate_args: bool = True                  # 可选
)

捷径 :对于明确的多分类任务,推荐使用 torchmetrics.classification.MulticlassConfusionMatrix(num_classes=10),参数更简洁,不需要手动指定 task


📝 三、参数详解

参数 类型 必填 默认值 说明
preds Tensor -- 模型预测。可以是概率/logits(浮点型)类别索引(整型) 。 多分类时,若是概率,形状通常为 (N, C);若是类别索引,形状为 (N,)
target Tensor -- 真实标签。多分类时形状为 (N,) 的整数张量。
task Literal["binary", "multiclass", "multilabel"] -- 任务类型。决定混淆矩阵的维度和内部转换逻辑。
num_classes Optional[int] 多分类时必填 None 类别总数。对于 10 分类,必须设为 10
num_labels Optional[int] 多标签时必填 None 标签总数,多标签任务专用。
threshold float 可选 0.5 二分类或多标签时将概率转为二值预测的阈值,多分类下忽略。
normalize Optional Literal\["true","pred","all" ] 可选 None 归一化方式: • None:输出原始计数值(整数张量)。 • "true":按行归一化(每行之和为 1),即每个真实类别下预测的分布(召回率视角)。 • "pred":按列归一化(每列之和为 1),即每个预测类别中有多少来自真实类别(精确率视角)。 • "all":除以所有样本总数,矩阵所有元素之和为 1。
ignore_index Optional[int] 可选 None 指定一个类别索引,计算时将其忽略(该类的真实和预测都不会计入矩阵)。常用于忽略填充标签。
validate_args bool 可选 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) 会在更新状态后返回当前累积的混淆矩阵。

代码示例

python 复制代码
import torch
import torchmetrics

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),即每个真实类别下预测的分布(召回率视角)。

python 复制代码
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 基本生命周期

python 复制代码
from torchmetrics import ConfusionMatrix

# 初始化(10分类)
confmat = ConfusionMatrix(task='multiclass', num_classes=10).to('cuda')

# 累积多个 batch
for batch in val_loader:
    preds, target = batch
    confmat.update(preds, target)

# 获取最终混淆矩阵
cm = confmat.compute()
print(cm.shape)        # torch.Size([10, 10])
print(cm)

# 重置状态,为下一轮准备
confmat.reset()

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

python 复制代码
# 在训练循环内可以直接调用对象
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 对象,你可以直接保存或记录。

python 复制代码
import torch
import torchmetrics
import matplotlib.pyplot as plt

# 1. 初始化(10分类)
confmat = torchmetrics.classification.MulticlassConfusionMatrix(num_classes=10)

# 2. 模拟累积数据
for _ in range(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 中记录

python 复制代码
class MyModel(pl.LightningModule):
    def __init__(self):
        super().__init__()
        self.confmat = ConfusionMatrix(task='multiclass', num_classes=10)

    def validation_step(self, batch, batch_idx):
        preds, target = batch
        self.confmat.update(preds, target)

    def on_validation_epoch_end(self):
        cm = self.confmat.compute()
        # 生成可视化图像并记录(假设你使用 TensorBoardLogger)
        fig = self.confmat.plot()
        if self.logger and hasattr(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,将混淆矩阵加入评估工具链,你就能既看到全局准确率,又能深入到每个类别的具体表现,做到"知其然,也知其所以然"。

相关推荐
旅僧18 小时前
王树森老师强化学习--同声传译版3
python·深度学习
额恩6620 小时前
阶段一:Vue 2 单页应用基础
人工智能·深度学习·机器学习
大模型码小白21 小时前
Milvus 架构设计:从向量索引到分布式检索,AI 原生存储引擎的内部机制
java·大数据·人工智能·分布式·深度学习·milvus
phltxy21 小时前
LangChain文本向量与检索器实践
人工智能·深度学习·语言模型·langchain
phltxy1 天前
LangChain从模型输出到RAG数据管道实战
服务器·人工智能·深度学习·语言模型·langchain
嘿丨嘿1 天前
VLA 入门(六):VLA 如何进行强化学习后训练?
人工智能·python·深度学习·机器人
嘿丨嘿1 天前
VLA 入门(二、三):机器人动作到底是什么?从关节空间到 Action Chunk
深度学习·机器学习·机器人
薛定e的猫咪2 天前
【模型推理】深度学习模型部署核心:算子Lower拆解、计算图优化与PNNX源码解析
人工智能·深度学习
Helen_cai2 天前
HarmonyOS ArkTS 实战:实现一个校园门禁与访客预约应用
深度学习·华为·harmonyos