如何在Python中生成单张多分类混淆矩阵?
生成多分类混淆矩阵的实现方案
1. 导入依赖库与加载数据
先导入所需工具库,再加载你的真实标签和预测标签数据:
import pandas as pd from sklearn.metrics import confusion_matrix import seaborn as sns import matplotlib.pyplot as plt # 加载数据 df_true = pd.DataFrame({ "y_true": [0,0,1,1,0,2] }) df_pred = pd.DataFrame({ "y_pred": [0,1,2,0,1,2] })
2. 计算混淆矩阵
提取标签列并计算混淆矩阵,矩阵中的每个值代表对应真实类别与预测类别的样本数量:
y_true = df_true['y_true'] y_pred = df_pred['y_pred'] # 计算混淆矩阵 confusion_mat = confusion_matrix(y_true, y_pred)
3. 可视化混淆矩阵
用热力图可视化混淆矩阵,清晰展示各类别的预测与真实值对应情况:
# 设置绘图参数 sns.set(font_scale=1.2) plt.figure(figsize=(6, 4)) # 绘制热力图,显示样本计数 ax = sns.heatmap(confusion_mat, annot=True, fmt='d', cmap='Blues', xticklabels=['类别0', '类别1', '类别2'], yticklabels=['类别0', '类别1', '类别2']) # 设置轴标签与标题 ax.set_xlabel('预测标签') ax.set_ylabel('真实标签') ax.set_title('多分类混淆矩阵') # 展示图形 plt.show()
关键参数说明
annot=True:在矩阵单元格内显示具体的样本数量fmt='d':确保以整数格式显示样本计数xticklabels/yticklabels:可替换为你的实际类别名称,比如分类任务中的"猫""狗""鸟"等
内容的提问来源于stack exchange,提问作者juanmac
相关产品推荐
相关产品推荐

