【问题标题】:Using non-linear scale with Seaborn heatmap使用 Seaborn 热图的非线性比例
【发布时间】:2020-04-21 16:47:37
【问题描述】:

我正在尝试为下面的热图使用对数刻度。我需要一个 0-30 之间数字的热图,然后是另一个可能是错误的较大值的配色方案。

尝试了几种不同的方法,但仍然非常困惑。感谢您的帮助。

干杯!

这是我正在使用的当前脚本。

read_occupancy = pd.read_csv (r'C:\Users\holborm\Desktop\Visualisation\dataaxisplotstuff.csv')   #read the csv file (put 'r' before the path string to address any special characters, such as '\'). Don't forget to put the file name at the end of the path + ".csv"

df = DataFrame(read_occupancy)    # assign column names


#create time and detector name axis

sns.heatmap(df.set_index('Row Labels').T, cmap='magma', linecolor='white', linewidths=.05)
sns.clustermap(df.set_index('Row Labels').T, cmap='magma', linecolor='white', linewidths=.05)

根据问题/答案更新

import pandas as pd
import seaborn as sns
import matplotlib.pyplot as plt
import matplotlib.ticker as ticker
from matplotlib.colors import LogNorm
def mix_palette():
    palette = sns.color_palette("GnBu", 10)
    palette[9] = sns.color_palette("OrRd", 10)[9]
    return palette


def set_ax(iax):
    for text in iax.texts:
        if float(text.get_text()) < 30:
            text.set_text("")
    iax.figure.tight_layout()


def load_data(path):
    initial = pd.read_csv(path, delim_whitespace=True)
    columns = list(initial.columns.values)[1:]
    rows = []
    for values in initial.values:
        rng = values[0]
        for column, value in zip(columns, values[1:]):
            rows.append([rng, column, value])
    return pd.DataFrame(data=rows, columns=['range', 'label', 'quantity'])
data = load_data('dataaxisplotstuff.csv')
data = data.pivot("range", "label", "quantity")
mi, ma = data.values.min(), data.values.max()
ax = sns.heatmap(data, cmap=mix_palette(), annot=True, square=True, cbar_kws={'ticks': ticker.LogLocator(numticks=8)},
                 xticklabels=True, yticklabels=True, norm=LogNorm(vmin=mi, vmax=ma))
set_ax(ax)
plt.show()

收到此错误

TypeError                                 Traceback (most recent call last)
<ipython-input-5-7466da1cd6c9> in <module>()
      1 data = load_data('dataaxisplotstuff.csv')
      2 data = data.pivot("range", "label", "quantity")
----> 3 mi, ma = data.values.min(), data.values.max()
      4 ax = sns.heatmap(data, cmap=mix_palette(), annot=True, square=True, cbar_kws={'ticks': ticker.LogLocator(numticks=8)},
      5                  xticklabels=True, yticklabels=True, norm=LogNorm(vmin=mi, vmax=ma))

~\AppData\Local\Continuum\anaconda3\lib\site-packages\numpy\core\_methods.py in _amin(a, axis, out, keepdims)
     27 
     28 def _amin(a, axis=None, out=None, keepdims=False):
---> 29     return umr_minimum(a, axis, None, out, keepdims)
     30 
     31 def _sum(a, axis=None, dtype=None, out=None, keepdims=False):

TypeError: '<=' not supported between instances of 'float' and 'str'

【问题讨论】:

  • 您能否提供指向“dataaxisplotstuff.csv”或其他格式相同的文件的链接,我加入了第一列 00 - 01hr => 00-01hr 和行标签 => RowLabels 的字符串。我猜你的错误可能与此有关。

标签: python seaborn


【解决方案1】:

我会试一试的。据我了解,您需要一个热图,其中 正常值 具有配色方案,异常值 具有不同颜色,热图也必须采用对数刻度。为此,我将使用pandas、seaborn 和matplotlib。版本为pandas: 0.22.0、matplotlib: 2.2.2 和seaborn: 0.9.0。首先是一些功能:

import matplotlib.pyplot as plt
import pandas as pd
import seaborn as sns
from matplotlib.colors import LogNorm


def mix_palette():
    palette = sns.color_palette("GnBu", 10)
    palette[9] = sns.color_palette("OrRd", 10)[9]
    return palette


def set_ax(iax):
    iax.collections[0].colorbar.set_ticklabels(['10', '30'])
    for text in iax.texts:
        if float(text.get_text()) < 30:
            text.set_text("")
    iax.figure.tight_layout()


def load_data(path):
    initial = pd.read_csv(path, delim_whitespace=True)
    columns = list(initial.columns.values)[1:]
    rows = []
    for values in initial.values:
        rng = values[0]
        for column, value in zip(columns, values[1:]):
            rows.append([rng, column, value])
    return pd.DataFrame(data=rows, columns=['range', 'label', 'quantity'])

函数mix_palette 创建一个混合调色板,set_ax 对图形进行一些调整,最后load_data 接收到一个指向 csv 的路径,就像示例中的路径一样,(使用空格作为分隔符) . load_data 的输出是 DataFrame,其形状与 seaborn 数据集中的航班相同,例如 (row_name, column_name, value)。现在绘图代码:

data = load_data('data.csv')
data = data.pivot("range", "label", "quantity")
mi, ma = data.values.min(), data.values.max()
ax = sns.heatmap(data, cmap=mix_palette(), annot=True, square=True, cbar_kws={'ticks': [10, 30],
                 xticklabels=True, yticklabels=True, norm=LogNorm(vmin=mi, vmax=ma))
set_ax(ax)
plt.savefig('image.png', bbox_inches='tight')
plt.show()

输出是: 这会将接近或高于 30 的值绘制为红色,并显示数值以便更好地可视化。更详细:

  • mix_palette 从默认调色板 "GnBu" 和 "OrRd" 创建一个混合。
  • 第一行set_ax 将颜色条(侧面的条)的标签设置为10 和30,循环将那些低于30 的单元格的值设置为空字符串。最后使布局紧凑(轴值的标签很大,您可以这样做以显示所有标签)。
  • cmap 参数接收调色板,annot=True 显示单元格的值,square=True 使热图的单元格为正方形,'ticks': [10, 30] 设置颜色条一侧的刻度位置,@ 987654345@ 是处理对数刻度的那个。
  • 要保存绘图,您可以使用该函数 plt.savefig('image.png', bbox_inches='tight') 确保在显示图像之前使用它。

【讨论】:

  • 太棒了!我正在研究的唯一另一件事是如何用指示值而不是 10-1 等标记热图
  • 什么是指示性值?
  • 看起来很完美 - 如何使尺寸变大?
  • 你可以删除 iax.figure.tight_layout() 行,但这会切断标签,我建议你看看这个answer
  • 我现在正试图把它变成一个可以反复用于其他类型数据的函数!
猜你喜欢
  • 2021-05-07
  • 2017-08-18
  • 2016-12-14
  • 2013-06-28
  • 2017-08-14
  • 2020-10-12
  • 1970-01-01
  • 2019-10-03
  • 1970-01-01
相关资源
最近更新 更多