【问题标题】:How to know scikit-learn confusion matrix's label order and change it如何知道 scikit-learn 混淆矩阵标签顺序并更改它
【发布时间】:2020-12-18 07:47:54
【问题描述】:

27个类存在多分类问题。

y_predict=[0 0 0 20 26 21 21 26 ....]

y_true=[1 10 10 20 26 21 18 26 ...]  

名为“answer_vocabulary”的列表存储了每个索引对应的 27 个单词。 answer_vocabulary=[0 1 10 11 2 3 农商东住北.....]

cm = 混淆矩阵(y_true=y_true, y_pred=y_predict)

我对混淆矩阵的顺序感到困惑。它是按升序排列的吗?如果我想用标签序列= [0 1 2 3 10 11农业商业生活东北...]重新排序混淆矩阵,我该如何实现它?

这是我尝试绘制混淆矩阵的函数。

def plot_confusion_matrix(cm, classes,
                        normalize=False,
                        title='Confusion matrix',
                        cmap=plt.cm.Blues):
    """
    This function prints and plots the confusion matrix.
    Normalization can be applied by setting `normalize=True`.
    """
    plt.imshow(cm, interpolation='nearest', cmap=cmap)
    plt.title(title)
    plt.colorbar()
    tick_marks = np.arange(len(classes))
    plt.xticks(tick_marks, classes, rotation=45)
    plt.yticks(tick_marks, classes)

    if normalize:
        cm = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis]
        print("Normalized confusion matrix")
    else:
        print('Confusion matrix, without normalization')

    print(cm)

    thresh = cm.max() / 2.
    for i, j in itertools.product(range(cm.shape[0]), range(cm.shape[1])):
        plt.text(j, i, cm[i, j],
            horizontalalignment="center",
            color="white" if cm[i, j] > thresh else "black")

    plt.tight_layout()
    plt.ylabel('True label')
    plt.xlabel('Predicted label')

【问题讨论】:

    标签: machine-learning scikit-learn deep-learning confusion-matrix


    【解决方案1】:

    来自 sklearn 的混淆矩阵不存储有关如何创建矩阵的信息(类排序和规范化):这意味着 您必须在创建后立即使用混淆矩阵 否则信息将丢失。

    默认情况下,sklearn.metrics.confusion_matrix(y_true,y_pred) 按照类在 y_true 中出现的顺序创建矩阵。

    如果您将此数据传递给 sklearn.metrix.confusion_matrix:

    +--------+--------+
    | y_true | y_pred |
    +--------+--------+
    | A      | B      |
    | C      | C      |
    | D      | B      |
    | B      | A      |
    +--------+--------+
    

    Scikit-leart 将创建这个混淆矩阵(省略零):

    +-----------+---+---+---+---+
    | true\pred | A | C | D | B | 
    +-----------+---+---+---+---+
    | A         |   |   |   | 1 |
    | C         |   | 1 |   |   |
    | D         |   |   |   | 1 |
    | B         | 1 |   |   |   |
    +-----------+---+---+---+---+
    

    它会返回这个 numpy 矩阵给你:

    +---+---+---+---+
    | 0 | 0 | 0 | 1 |
    | 0 | 0 | 1 | 0 |
    | 0 | 0 | 0 | 1 |
    | 1 | 0 | 0 | 0 |
    +---+---+---+---+
    

    如果您想选择类或重新排序它们,您可以将“标签”参数传递给confusion_matrix()。

    重新排序:

    labels = ['D','C','B','A']
    mat = confusion_matrix(true_y,pred_y, labels=labels)
    
    

    或者,如果您只想关注一些标签(如果您有很多标签,这很有用):

    labels = ['A','D']
    mat = confusion_matrix(true_y,pred_y, labels=labels)
    

    另外,看看sklearn.metrics.plot_confusion_matrix。它适用于小型 (

    如果您有 >100 个类,则绘制矩阵需要白色。

    【讨论】:

      【解决方案2】:

      生成的混淆矩阵中的列/行的顺序与sklearn.utils.unique_labels() 返回的相同,后者提取“唯一标签的有序数组”。在confusion_matrix()(main, git-hash 7e197fd)的source code中,感兴趣的行如下

      if labels is None:
          labels = unique_labels(y_true, y_pred)
      else:
          labels = np.asarray(labels)
      

      这里,labels 是 confusion_matrix() 的可选参数,用于自行规定标签的排序/子集:

      cm = confusion_matrix(true_y, pred_y, labels=labels)
      

      因此,如果labels = [0, 10, 3],cm 将具有形状(3,3),行/列可以直接用labels 进行索引。如果你知道熊猫:

      import pandas as pd
      cm = pd.DataFrame(cm, index=labels, columns=labels)
      

      请注意,unique_labels() 的文档声明不支持混合类型的标签(数字和字符串)。在这种情况下,我建议使用LabelEncoder。这将使您免于维护自己的查找表。

      from sklearn.preprocessing import LabelEncoder
      encoder = LabelEncoder()
      y = encoder.fit_transform(y)
      
      # y have now values between 0 and n_labels-1.
      # Do some ops here...
      ...
      
      # To convert back:
      y_pred = encoder.inverse_transform(y_pred)
      y = encoder.inverse_transform(y)
      

      正如previous answer 已经提到的,plot_confusion_matrix() 可以方便地可视化混淆矩阵。

      【讨论】:

        猜你喜欢
        • 2016-05-12
        • 2020-05-18
        • 2020-05-30
        • 2014-06-11
        • 2018-05-22
        • 2018-01-15
        • 2018-10-23
        • 2020-07-16
        • 2020-01-22
        相关资源
        最近更新 更多