torchmetrics堪称模型评估界的"绝世秘籍",招式精妙且威力无穷。若想真正参透其中玄机、融会贯通,列位看官莫急,且听我细细拆解。这是 torchmetrics 系列文章的第二篇。
聚焦在 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)
👀 具体怎么看?
- 横向看一行(关注"真实")→ 看召回率 (Recall)
问题: "所有真实的猫,模型找全了吗?"
- 怎么看:盯着**"猫"的那一行**。
- 含义:这一行代表世界上所有真实的猫。对角线上的数字是找对的,非对角线上的数字是漏掉的(被误判成了狗或鸟)。
- 用途:检查模型是否漏掉了某个类别的样本。
- 竖向看一列(关注"预测")→ 看精确率 (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方法直观观察。
结合之前的 Accuracy 和 Macro-F1,将混淆矩阵加入评估工具链,你就能既看到全局准确率,又能深入到每个类别的具体表现,做到"知其然,也知其所以然"。