【问题标题】:Ticks position in heatmap with categorical data (seaborn)带有分类数据的热图中的刻度位置(seaborn)
【发布时间】:2020-05-30 17:38:37
【问题描述】:

我正在尝试绘制我的预测的混淆矩阵。我的数据是多类的(13 个不同的标签),所以我使用的是热图。

正如您在下面看到的,我的热图看起来一般都不错,但标签位置有点不对:y 刻度应该稍微低一些,x 刻度应该更靠右一些。我想稍微移动两个轴刻度,以便它们与每个正方形的中心对齐。

我的代码:

sns.set()
my_mask = np.zeros((con_matrix.shape[0], con_matrix.shape[0]), dtype=int)
for i in range(con_matrix.shape[0]):
    for j in range(con_matrix.shape[0]):
        my_mask[i][j] = con_matrix[i][j] == 0 

fig_dims = (10, 10)
plt.subplots(figsize=fig_dims)
ax = sns.heatmap(con_matrix, annot=True, fmt="d", linewidths=.5, cmap="Pastel1", cbar=False, mask=my_mask, vmax=15)
plt.xticks(range(len(party_names)), party_names, rotation=45)
plt.yticks(range(len(party_names)), party_names, rotation='horizontal')
plt.show()

为了复制目的,这里是con_matrixparty_names 硬编码:

import numpy as np
from matplotlib import pyplot as plt
import seaborn as sns 

con_matrix = np.array([[55, 0, 0, 0,0, 0, 0,0,0,0,0,0,2], [0,199,0,0,0,0,0,0,0,0,2,0,1],
 [0, 0,52,0,0,0,0,0,0,0,0,0,1],
 [0,0,0,39,0,0,0,0,0,0,0,0,0],
 [0,0,0,0,90,0,0,0,0,0,0,4,3],
 [0,0,0,1,0,35,0,0,0,0,0,0,0],
 [0,0,0,0,5,0,26,0,0,1,0,1,0],
 [0,5,0,0,0,1,0,44,0,0,3,0,1],
 [0,1,0,0,0,0,0,0,52,0,0,0,0],
 [0,1,0,0,2,0,0,0,0,235,0,1,1],
 [1,2,0,0,0,0,0,3,0,0,34,0,3],
 [0,0,0,0,5,0,0,0,0,1,0,40,0],
 [0,0,0,0,0,0,0,0,0,1,0,0,46]])

party_names = ['Blues', 'Browns', 'Greens', 'Greys', 'Khakis', 'Oranges', 'Pinks', 'Purples', 'Reds', 'Turquoises', 'Violets', 'Whites', 'Yellows']

我已经尝试使用不同轴的position 参数,但结果并不好。在这个网站上也找不到确切的答案(至少不是适用于分类数据的解决方案)。

我是使用 seaborn 进行可视化的新手,如有任何改进,我们将不胜感激(不仅针对我的问题,还针对我的代码和可视化)。

【问题讨论】:

  • 只需使用 plt.xticks(np.arange(0.5, len(party_names)), ... 和类似的 y。这样,刻度就可以很好地定位在每个单元格的中心。
  • 请注意,您也可以直接添加刻度标签:ax = sns.heatmap(..., xticklabels=party_names, yticklabels=party_names)。您仍然可以使用plt.xticks(rotation=45) 旋转它们而无需额外的参数。

标签: matplotlib data-visualization seaborn heatmap


【解决方案1】:

您可以将两个刻度标签移动 0.5 个偏移量以获得所需的对齐方式。为此,我使用了 NumPy 的 arange,它可以将 0.5 向量化加到整个数组中。

plt.xticks(np.arange(len(party_names))+0.5, party_names, rotation=45)
plt.yticks(np.arange(len(party_names))+0.5, party_names, rotation='horizontal')

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2017-04-05
    • 2020-12-29
    • 2021-04-16
    • 2018-05-26
    • 1970-01-01
    相关资源
    最近更新 更多