做分类模型的朋友,应该都遇过这种场景:训练完一看准确率,97%,心里美滋滋;结果把模型拿到业务方一跑,对方来一句“你这模型跟瞎猜有什么区别”,自己还一脸懵。后来我才明白,问题多半出在我只盯着准确率,没有去看混淆矩阵。
这篇博文不绕弯子,专门讲 sklearn 里的 confusion_matrix 函数怎么理解、怎么用、怎么避坑。我会从二分类的四格矩阵讲到多分类的矩阵解读和可视化,再把我平时踩过的坑一并列出来。适合刚接触 sklearn 的初学者,也适合已经在调模型但总觉得评估报告差点意思的工程师。看完你就能自己写出一份能说服业务方的模型评估报告。
1. 为什么有准确率还不够:混淆矩阵到底解决了什么问题
1.1 一个97%准确率的模型,为什么被业务方嫌弃
先举个极端但真实的例子。假设线上有 100 个样本,95 个是正常用户,5 个是风险用户。你训练了一个模型,为了省事它学会了“不管什么样本,一律预测成正常用户”。这样准确率是 95/100 = 95%,听起来很棒对吧?但风险用户一个都没捞出来——这一类的召回率是 0。
如果这个场景是欺诈交易检测,那这个模型上线就是在帮倒忙。准确率高不代表模型好,它只是看“整体”正确率,会把少数类的灾难掩盖掉。混淆矩阵则完全不同,它把所有真实类别和预测类别的组合都摊开给你看,一目了然:哪些类被照顾得好,哪些类被彻底放弃。
正是因为这种透明度,我在实际项目里几乎不会只丢一个 accuracy 出去汇报。无论是验证集还是线上评估,第一件事都是先跑混淆矩阵,再去看派生出来的指标。混淆矩阵相当于模型的“体检报告”,能看到整体健康指标背后的具体科室问题。
1.2 二分类里的四个格子:TN、FP、FN、TP分别代表什么
先用一个小的手写数据集感受一下混淆矩阵的形状和含义。比如我们有 8 个样本,真实标签和预测标签长这样:
from sklearn.metrics import confusion_matrix y_true = [0, 0, 0, 1, 1, 1, 1, 1] y_pred = [0, 0, 1, 0, 1, 1, 1, 1] cm = confusion_matrix(y_true, y_pred) print(cm)输出是:
[[2 1] [1 4]]这个 2x2 矩阵怎么读?关键记住一句话:行是真实类别,列是预测类别。也就是说:
- 第一行代表真实为 0 的样本,一共 3 个,其中 2 个被预测成 0,1 个被预测成 1。
- 第二行代表真实为 1 的样本,一共 5 个,其中 1 个被预测成 0,4 个被预测成 1。
用正式术语来命名这四个格子:
- cm[0][0] = TN(True Negative),真实为负且预测为负,正确拒识;
- cm[0][1] = FP(False Positive),真实为负但预测为正,误报;
- cm[1][0] = FN(False Negative),真实为正但预测为负,漏报;
- cm[1][1] = TP(True Positive),真实为正且预测为正,正确命中。
这里要特别提醒一点:sklearn 默认的矩阵布局和很多教科书上的 [[TP, FP], [FN, TP]] 不一样。sklearn 是按标签数值从小到大排列的,所以第一行是标签 0,第二行是标签 1。如果你用惯了“左上角是TP”的写法,第一次看 sklearn 的矩阵会觉得是反的,这是新手最容易懵的地方。
1.3 从四格矩阵算出准确率、精确率、召回率与F1
矩阵本身不是终点,它的价值在于能派生出各种评估指标。基于上面这个矩阵:
- 准确率 Accuracy = (TP + TN) / 总和 = (4 + 2) / 8 = 0.75;
- 精确率 Precision = TP / (TP + FP) = 4 / (4 + 1) = 0.8,预测为正的样本里真正是正类的比例;
- 召回率 Recall = TP / (TP + FN) = 4 / (4 + 1) = 0.8,真实正类中被找回来的比例;
- F1 = 2 * Precision * Recall / (Precision + Recall) = 0.8。
算完你再用 sklearn 的 classification_report 对照一下:
from sklearn.metrics import classification_report print(classification_report(y_true, y_pred, labels=[0, 1]))它会把类别 0 和类别 1 各自的精确率、召回率都输出一遍。注意:在多分类视角下,每个类别都可以看成“正类”,所以“精确率”这个概念是针对某一个类来讲的。这也解释了为什么工程上要同时看分类报告和混淆矩阵,而不是只看一个总准确率。
2. confusion_matrix函数:参数拆解与三个容易踩坑的细节
2.1 函数签名与核心参数速览
confusion_matrix 的函数签名如下:
sklearn.metrics.confusion_matrix( y_true, y_pred, *, labels=None, sample_weight=None, normalize=None, )日常使用中最常碰到的就是 y_true、y_pred、labels 和 normalize 四个参数。y_true 是真实标签,y_pred 是模型预测标签,这两个非常好理解。容易忽略的是 labels 和 normalize,这两兄弟一个是“矩阵行列顺序”的控制器,一个是“矩阵数值展示形式”的开关。如果把这两个玩明白,你在很多实际场景里都能少踩坑。
下面这个表格可以帮你快速对参数功能建立印象:
| 参数 | 作用 | 注意事项 |
|---|---|---|
| y_true | 真实标签数组 | 必须是一维数组或列表 |
| y_pred | 模型预测标签数组 | 长度必须与 y_true 一致 |
| labels | 指定矩阵行列的标签列表 | 不指定时默认按并集排序,容易出幺蛾子 |
| sample_weight | 样本权重 | 权重存在时矩阵输出的是加权计数 |
| normalize | 归一化方式 | 可选 'true'、'pred'、'all',默认 None |
2.2 labels参数:不显式指定,矩阵顺序会“偷偷”变
很多人刚用这个函数时只传两个必填参数,遇到矩阵行列顺序不是自己预期的情况就开始懵。根因基本都在 labels 上。不传 labels 时,sklearn 内部会对 y_true 和 y_pred 的并集做排序,再按这个顺序生成矩阵。也就是说,如果模型在测试集上某种类别完全没有出现,默认 labels 就不会含有它,矩阵的行列也就跟着少一块。
举个实际例子:你训练的时候有 [0, 1, 2] 三类,但测试集里恰好没有类别 2 的真实样本,模型也没预测出类别 2,默认矩阵就只有 2x2 了。这样不仅展示上莫名其妙,后续算每个类别的召回率还会漏掉一类。我吃过这个亏后,现在写代码一律显式传 labels,最稳妥的方式是直接用训练好的模型的model.classes_,或者用你预先定义好的类别列表:
labels = ['setosa', 'versicolor', 'virginica'] cm = confusion_matrix(y_true, y_pred, labels=labels)还有一个细节:如果标签本身是字符串,比如['cat', 'dog', 'bird'],labels 参数也要传对应字符串列表。保持 y_true、y_pred 和 labels 三者的取值语义一致,矩阵才不会出现值对不上的离谱错误。
2.3 normalize参数:三种归一化方式的适用场景
normalize 参数用来把矩阵里的绝对计数变成比例,好处是不同数据集之间可以直接对比,尤其在做模型报告时更有说服力。三种模式分别解决三个问题:
| normalize 取值 | 计算方式 | 解决什么问题 |
|---|---|---|
'true' | 每一行除以该行总和 | 看真实类别被预测到各个类别的比例,定位召回率短板 |
'pred' | 每一列除以该列总和 | 看预测类别的来源分布,定位精确率短板 |
'all' | 所有格子除以总数 | 看各组合占全体的比例,直观展示整体错误分布 |
拿前面那个八样本的例子来说:
print(confusion_matrix(y_true, y_pred, normalize='true')) # [[0.66666667 0.33333333] # [0.2 0.8 ]] print(confusion_matrix(y_true, y_pred, normalize='pred')) # [[0.66666667 0.2 ] # [0.33333333 0.8 ]] print(confusion_matrix(y_true, y_pred, normalize='all')) # [[0.25 0.125] # [0.125 0.5 ]]注意一下,normalize='true' 时每一行和为 1,normalize='pred' 时每一列和为 1,normalize='all' 时全部和为 1。我在类别不平衡的数据集上,最喜欢用 normalize='true'。因为少数类样本本来就少,看绝对数值容易造成“类别 A 有 50 个错误,类别 B 只有 5 个错误”的错觉,行归一化之后就看到比例,更公平。
这里还要提示一下浮点数的问题:normalize 之后打印出来的矩阵可能会出现0.30000000000000004这种小数尾巴,这是浮点数运算的正常现象,不是函数算错了。展示的时候记得格式化,比如保留两位小数。
2.4 sample_weight参数:什么时候需要给样本加权
sample_weight 用得不像 labels 和 normalize 那么频繁,但它很实用。这个参数的作用是给每个样本一个权重,矩阵计算时把“样本数量”替换成“样本权重之和”。常见场景有两个:一是你在做数据修正,某些样本来源更可靠,给它们更高权重;二是你在做上采样或下采样之后,想用一个“伪样本数”来评估效果。
比如你有一批样本,重复抽样了 3 次,如果不改权重直接算混淆矩阵,相当于每条记录被平等对待。但如果你把重复样本的 weight 设成 1/3,矩阵算出来的就是去重后的加权口径。这个技巧在评估链路里非常有用。注意一旦用了 sample_weight,矩阵内元素不一定是整数,直接画图或者读数字时要有心理准备。
3. 多分类场景下的完整实操:从矩阵解读到可视化
3.1 多分类矩阵的正确打开方式:先看行,再看列
二分类有四格术语,多分类就不再单独叫 TN、FP 这些了,泛化为“真实第 i 类、预测第 j 类”的组合。多分类矩阵的解读技巧可以总结为一句话:先看行,再看列,最后盯非对角线。
- 每一行的数值总和等于该真实类别在测试集里的样本总数;
- 每一列的数值总和等于模型预测为该类别的样本总数;
- 对角线是预测正确的样本,非对角线全是错误;
- 如果某一行对角线之外的数值很大,说明这类样本经常被分错到别的类;
- 如果某一列对角线之外的数值很大,说明模型经常把别的类预测成这个类。
这种“行视角”和“列视角”的差异,能帮你快速定位模型到底是在“漏”(行视角召回率低)还是在“多报”(列视角精确率低)。比如在一个三类任务里,模型经常把类 1 和类 2 搞混。如果你发现真实类 1 的样本有一大批被预测成类 2,说明模型对类 1 的判别能力不足;如果你发现预测类 2 里混了很多真实类 1 的样本,说明类 2 的判定边界太宽了。
3.2 一次完整实操:鸢尾花数据 + SVM分类器
理论说得再多,不如跑一遍。这里我用 sklearn 自带的鸢尾花数据集,取前两个特征(花萼长宽)来做分类。使用前两个特征的前提就是“模型一定会有错分”,这样我们才能看到非对角线上的“混淆”效果,比全对的矩阵有讲解价值。
from sklearn.datasets import load_iris from sklearn.model_selection import train_test_split from sklearn.svm import SVC from sklearn.metrics import confusion_matrix, classification_report # 加载数据并只取前两个特征,故意制造容易混淆的效果 data = load_iris() X = data.data[:, :2] y = data.target X_train, X_test, y_train, y_test = train_test_split( X, y, test_size=0.3, random_state=42, stratify=y ) model = SVC(kernel='rbf', C=1.0, gamma='scale', random_state=42) model.fit(X_train, y_train) y_pred = model.predict(X_test) labels = [0, 1, 2] cm = confusion_matrix(y_test, y_pred, labels=labels) print("混淆矩阵:") print(cm) print() report = classification_report( y_test, y_pred, labels=labels, target_names=data.target_names ) print("分类报告:") print(report)在我本地的输出大致是:
混淆矩阵: [[15 0 0] [ 1 12 2] [ 0 3 12]] 分类报告: precision recall f1-score support setosa 1.00 1.00 1.00 15 versicolor 0.80 0.80 0.80 15 virginica 0.86 0.80 0.82 15 accuracy 0.87 45 macro avg 0.89 0.87 0.87 45 weighted avg 0.89 0.87 0.87 45这个矩阵看着就很有信息量了。setosa(山鸢尾)表现完美,15 个全中;versicolor(变色鸢尾)有 1 个被错判成 setosa,2 个被错判成 virginica;virginica(维吉尼亚鸢尾)有 3 个被错判成 versicolor。如果你只看 0.87 的 accuracy,只会得到一个“不错”的结论。但通过矩阵你能立刻发现模型的主要软肋集中在 versicolor 和 virginica 之间——这是因为只用花萼长宽两个特征时,这两类的特征分布重叠比较严重。
这就是混淆矩阵的实战价值:它能告诉你错误集中发生在“哪些类之间”,而不只是告诉你“错得很多”。
3.3 两种可视化方案:ConfusionMatrixDisplay与seaborn热力图
纯数字矩阵在老练的工程师手里已经够用,但汇报给业务方或写进文档时,一张带颜色深浅的图远比数字有冲击力。这里我分享两种常用方案。
方案一是 sklearn 自带的 ConfusionMatrixDisplay,优点是代码短、零依赖(只需要 matplotlib):
import matplotlib.pyplot as plt from sklearn.metrics import ConfusionMatrixDisplay disp = ConfusionMatrixDisplay.from_predictions( y_test, y_pred, display_labels=data.target_names, cmap='Blues', ) plt.show()如果分类器已经训练好,也可以用 from_estimator 直接传入模型和数据集特征,让它内部自己预测并绘制,更加省事。需要注意:from_predictions 和 from_estimator 返回的对象自带一个 ax,如果你想把矩阵画到你自己的子图网格里,可以传 ax 参数进去。
方案二是用 seaborn 的 heatmap 手动绘制,自由度更高,控制力更强。我平时做正式报告更偏爱这种,因为可以微调字体、边框、配色和保存精度:
import matplotlib.pyplot as plt import seaborn as sns plt.figure(figsize=(6, 5)) sns.heatmap( cm, annot=True, fmt='d', cmap='Blues', xticklabels=data.target_names, yticklabels=data.target_names, cbar=True, linewidths=0.5, linecolor='gray', annot_kws={"size": 12}, ) plt.xlabel('预测标签') plt.ylabel('真实标签') plt.tight_layout() plt.savefig('confusion_matrix.png', dpi=150, bbox_inches='tight') plt.show()fmt='d'表示显示整数,如果要显示百分比或小数,改成fmt='.2f'。这里要特别提醒一个中文显示问题:如果 xlabel、ylabel 或刻度标签用了中文,matplotlib 默认字体可能显示成方框。稳妥做法是直接用英文标签,或者先配置系统支持的中文字体再绘图。
另外,如果想让图更直观地表达“每类召回率”,可以先算好 normalize='true' 的矩阵,再传给 heatmap。这样每个格子的数值代表“真实类别被预测成各类的比例”,对角线越接近 1 越好。我在模型迭代对比时经常用这种带比例的热力图,视觉上比绝对计数更公平。
4. 高频坑位与实战经验总结
4.1 常见问题速查表
下面这张表是我在实际使用中比较容易踩的坑,这几点理解了能少走很多弯路。
| 问题现象 | 根本原因 | 解决办法 |
|---|---|---|
| 矩阵看起来是“反”的 | 教材多用 [[TP, FP], [FN, TP]] 布局,sklearn 按 labels 排序 | 记住行=真实,列=预测,先看 labels 确认顺序 |
| 类别顺序跟想象不一致 | 没有显式传 labels,默认按并集排序 | 显式传 labels,最好用 model.classes_ 或预定义类别表 |
| 多分类矩阵缺少某一行/列 | 该类别在测试数据和预测结果中都未出现 | 用完整类别列表显式指定 labels |
| 归一化后出现 0.30000000000000004 | 浮点运算精度问题 | 展示时用 round 或格式化字符串 |
| 某列全为 0 | 模型完全没有预测出该类别 | 检查数据分布和模型是否失效,可能是训练集缺失某类 |
| 中文标签画图变方框 | matplotlib 默认字体不支持中文 | 配置中文字体,或使用英文标签 |
这里额外说一个最常见的“低级但致命”错误:y_true 和 y_pred 长度不一致。confusion_matrix 不会给你太友好的提示,甚至会因为广播机制算出奇怪的矩阵。所以每次调用前我都先确认一下len(y_true) == len(y_pred),养成习惯后能避免很多莫名其妙的结果。
4.2 样本不均衡下的矩阵读法
很多分类问题都有类别不平衡,比如信用卡欺诈、工单故障、异常检测,少数类可能只占 1%-5%。这时看绝对计数的矩阵基本看不出东西,因为少数类的支持度太小,容易被淹没在多数类的大数字里。
我的实操经验是:矩阵绝对计数配合 normalize='true' 的矩阵一起看。绝对计数回答“有多少”,归一化回答“占多大比例”。比如 100 个测试样本里欺诈只有 1 个,模型把 99 个正常样本全预测对了,但唯一那个欺诈样本也被预测成正常。绝对计数矩阵长这样:
[[99 0] [ 1 0]]95% 的准确率看起来很高,但第二行的召回率是 0%。这时候 banlance 一下视角,你才会发现模型本质上是个“只会说正常”的哑巴。所以我在不平衡场景下很少提 accuracy,而是盯着少数类的 recall 和曲线下面积,矩阵只是辅助确认错误分布的工具。
4.3 我总结的几个工程化建议
第一,把混淆矩阵的调用封装成工具函数。项目里别到处裸写confusion_matrix(y_test, y_pred),而是封装一个report_model(y_true, y_pred, labels)之类的函数,内部同时输出矩阵、分类报告和归一化矩阵,一次调用全出来。这样不同人写的代码风格也能统一,减少沟通成本。
第二,训练前就把类别字典定义清楚。不管你是二分类还是多分类,先定义一个LABELS = ['class_a', 'class_b', ...]的常量,后续所有调用都用它。这种方式看起来多写一行,但能避免测试集里某个类别缺失时矩阵突然变形的尴尬。
第三,矩阵和指标要配套使用。混淆矩阵偏“诊断”,classification_report 偏“评分”。诊断告诉你错在哪里,评分告诉你严重程度。只给矩阵不给报告,读者看不出整体水平;只给报告不给矩阵,又不知道具体错误模式。两个一起上,报告才完整。
我自己在实际项目中的习惯是:模型训练完第一件事不是看准确率,而是先打印confusion_matrix(y_test, y_pred, labels=model.classes_, normalize='true')。因为这个矩阵能直接告诉我模型在哪个类上偷了懒。另外一个很实用的小技巧是:保存矩阵图时,把 display_labels 传成业务上能看懂的名称,而不是 0、1、2 这种数字。很多业务方不关心指标公式,但他们看得懂“正常品被误判成次品”“A类错标成B类”这种直白的矩阵图,汇报效率立刻提升。