【问题标题】:Plotting already calculated Confusion Matrix using Python使用 Python 绘制已经计算好的混淆矩阵
【发布时间】:2020-06-14 06:32:51
【问题描述】:

对于已经给定的混淆矩阵值,我如何在 Python 中绘制类似于 here 所示的混淆矩阵?

在代码中,他们使用 sklearn.metrics.plot_confusion_matrix 方法,该方法根据基本事实和预测计算混淆矩阵。

但就我而言,我已经计算了我的混淆矩阵。例如,我的混淆矩阵是(百分比值):

[[0.612, 0.388]
 [0.228, 0.772]]

【问题讨论】:

标签: python matplotlib confusion-matrix


【解决方案1】:

我使用 seaborn 的热图。你可以定义一个方法:

import numpy as np
import seaborn as sns; sns.set_theme()
sns.set(font_scale=2)

def plot_matrix(cm, classes, title):
  ax = sns.heatmap(cm, cmap="Blues", annot=True, xticklabels=classes, yticklabels=classes, cbar=False)
  ax.set(title=title, xlabel="predicted label", ylabel="true label")

并使用:

cm = np.array([[0.612, 0.388], [0.228, 0.772]])
classes = ['class A', 'class B']
title = "title example"

plot_matrix(cm, classes, title)

输出是这样的:

【讨论】:

    【解决方案2】:

    我看到有人已经回答了这个问题,但我正在添加一个对作者甚至其他用户有用的新问题。

    可以在 Python绘制一个已经通过mlxtend 包计算的混淆矩阵

    Mlxtend(机器学习扩展)是一个有用的 Python 库 日常数据科学任务的工具。

    片段代码:

    # Imports
    from mlxtend.plotting import plot_confusion_matrix
    import matplotlib.pyplot as plt
    import numpy as np
    
    # Your Confusion Matrix
    cm = np.array([[0.612, 0.388],
                   [0.228, 0.772]])
    
    # Classes
    classes = ['class A', 'class B']
    
    figure, ax = plot_confusion_matrix(conf_mat = cm,
                                       class_names = classes,
                                       show_absolute = False,
                                       show_normed = True,
                                       colorbar = True)
    
    plt.show()
    

    输出将是:

    【讨论】:

      【解决方案3】:

      如果您检查source 中的sklearn.metrics.plot_confusion_matrix,您可以看到如何处理数据以创建绘图。然后你可以重用构造函数ConfusionMatrixDisplay 并绘制你自己的混淆矩阵。

      import matplotlib.pyplot as plt
      from sklearn.metrics import ConfusionMatrixDisplay
      
      cm = [0.612, 0.388, 0.228, 0.772] # your confusion matrix
      ls = [0, 1] # your y labels
      disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=ls)
      disp.plot(include_values=include_values, cmap=cmap, ax=ax, xticks_rotation=xticks_rotation)
      plt.show()
      

      【讨论】:

      • 当我运行这个时,我收到关于未定义形状的错误。我将混淆矩阵行重写为cm = np.array([[tn,fp], [fn,tp]]),我将在其中转换为numpy数组。我还冒昧地创建了代表真假阳性和阴性的变量。
      猜你喜欢
      • 1970-01-01
      • 2016-01-31
      • 2018-04-01
      • 1970-01-01
      • 2017-02-25
      • 2020-07-15
      • 2021-07-21
      相关资源
      最近更新 更多