sklearn.confusion_matrix默认行是真实标签、列是预测标签,需确保整数标签对齐、显式指定labels参数,seaborn.heatmap中xticklabels对应预测、yticklabels对应真实,归一化用normalize='true'并配fmt='.2f'和cbar_kws,中文显示需设置字体及旋转标签。

直接用 sklearn.metrics.confusion_matrix + seaborn.heatmap 就能画出来,但多数人卡在标签对不齐、归一化失效、中文显示乱码、热力图数值错位这四步上。
sklearn.confusion_matrix 输出的矩阵为什么和 seaborn.heatmap 显示不一致?
根本原因是:默认输出是「行=真实标签,列=预测标签」,但很多人误以为是「行=预测,列=真实」。如果类名顺序没对齐,整个矩阵就翻转了。
- 确保
y_true和y_pred是整数标签(不是 one-hot),且取值范围是0到n_classes-1 - 传给
confusion_matrix时显式指定labels参数,例如labels=list(range(len(class_names))) -
seaborn.heatmap的xticklabels对应列(预测),yticklabels对应行(真实)——别反了 - 加
fmt='d'防止科学计数法,加annot=True才显示数字
归一化后 heatmap 数值全是 0.0 或 1.0?
因为 normalize='true' 是按行归一(每行和为 1),normalize='pred' 是按列归一,normalize='all' 是全局归一。多数人想看「每个类别的召回率」,该用 normalize='true',但必须配合 cmap 和 vmin/vmax 控制颜色范围,否则小数会全挤在浅色区。
图片提示词生成器?不止如此。 马甲系统 —— 把脑海中的画面,翻译成AI能理解的专业表达。 用得越多,它越懂你:首次需要多问几句确认方向,用久了几乎一说就懂。 用得越多,它越快:缓存机制让后续对话越来越省。 RAG进化:成功案例持续入库,越跑越聪明。 输入「新手指南」查看完整功能介绍
- 归一化后数据类型变成
float64,fmt='d'会报错,要改用fmt='.2f' - 加
cbar_kws={'format': '.2f'}让 colorbar 标签也匹配小数位 - 若想保留整数标签但显示百分比,可手动除以行和再 ×100,然后用
fmt='.0f'+ 加 % 符号(需自定义annot内容)
中文类别名显示为方块或重叠?
Matplotlib 默认字体不支持中文,且 xticklabels/yticklabels 太长时会自动重叠或截断。
- 开头加这两行解决字体问题:
plt.rcParams['font.sans-serif'] = ['SimHei', 'Arial Unicode MS'];plt.rcParams['axes.unicode_minus'] = False - 用
plt.xticks(rotation=30)和plt.yticks(rotation=0)避免横排拥挤 - 调整图像尺寸:
plt.figure(figsize=(len(class_names)*1.2, len(class_names)*1.0)),防止被裁剪 - 最后调用
plt.tight_layout(),不然标签可能被切掉
PyTorch/TensorFlow 模型预测后怎么喂给 confusion_matrix?
模型输出是 logits 或概率,不能直接塞进 confusion_matrix。必须先取 argmax 转成整数类别索引,且注意设备(GPU tensor 需先 .cpu())和维度(常有多余 batch 维度)。
- PyTorch 示例:
y_pred = model(x).argmax(dim=1).cpu().numpy();y_true = y.cpu().numpy() - TensorFlow/Keras 示例:
y_pred = model.predict(x_test).argmax(axis=1);y_true = y_test.argmax(axis=1)(若 y_test 是 one-hot) - 如果原始标签是字符串(如
['cat', 'dog']),得用LabelEncoder或字典映射回整数,且确保训练/测试用同一套映射 - 务必检查
y_true和y_pred长度是否一致,否则confusion_matrix报错ValueError: Found array with dim 3. Expected
最容易被忽略的是:混淆矩阵本身不处理类别缺失。如果某类在测试集中没出现,confusion_matrix 默认不会补零行/列,导致 shape 对不上 class_names。必须显式传 labels 参数强制对齐维度。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!










