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)
👀 具体怎么看?
- 横向看一行(关注“真实”)→ 看召回率 (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:float=0.5,# 可选,默认0.5,二分类/多标签用normalize:Optional[Literal["true","pred","all"]]=None,# 可选,默认None(输出整数计数)ignore_index:Optional[int]=None,# 可选,默认Nonevalidate_args:bool=True# 可选,默认True)->Tensor2.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。
📝 三、参数详解
| 参数 | 类型 | 必填 | 默认值 | 说明 |
|---|---|---|---|---|
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)会在更新状态后返回当前累积的混淆矩阵。
代码示例:
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方法直观观察。
结合之前的Accuracy和Macro-F1,将混淆矩阵加入评估工具链,你就能既看到全局准确率,又能深入到每个类别的具体表现,做到“知其然,也知其所以然”。