如何仅在热力图的下三角区域显示网格?
问题:热力图仅在下三角区域保留网格
我制作了一个遮罩上三角部分的热力图,但设置linewidths后网格会显示在整个热力图区域,希望仅在下三角区域保留网格。
原代码:
import matplotlib.pyplot as plt import pandas as pd import seaborn as sb import numpy as np import seaborn as sns plt.figure(figsize=(10,10)) # importing Dataset data_df = pd.read_csv("My data.csv") #corr matrix r=data.corr() r[abs(r)<0.00001]=0 r.round(decimals=4, out=None) params = {'mathtext.default': 'regular' } plt.rcParams.update(params) # mask mask=np.ones([7,7]) mask *= 1-np.tri(*mask.shape,k=0) # plotting a masked correlation heatmap dataplot = sb.heatmap(r, cmap="gist_heat", vmin=-0.2, vmax=1.1,annot=True, mask=mask,fmt='.4f',annot_kws={"fontsize":15},linewidths=1, linecolor='black',cbar_kws={"shrink": 0.81}, clip_on=False,square=True) plt.xticks(rotation=90) plt.yticks(rotation=0) plt.savefig('/content/heatmap.png',bbox_inches = 'tight',dpi=330) # displaying heatmap plt.show()
解决方案
Seaborn的heatmap默认会给所有单元格(包括被mask遮罩的区域)添加边框,所以需要手动控制仅在下三角区域绘制网格线。修改后的代码如下:
import matplotlib.pyplot as plt import pandas as pd import seaborn as sns import numpy as np plt.figure(figsize=(10,10)) # 导入数据 data_df = pd.read_csv("My data.csv") # 计算相关矩阵 r = data_df.corr() r[abs(r) < 0.00001] = 0 r = r.round(decimals=4) # 设置matplotlib参数 params = {'mathtext.default': 'regular' } plt.rcParams.update(params) # 上三角遮罩 mask = np.ones([7,7]) mask *= 1 - np.tri(*mask.shape, k=0) # 先绘制不带网格的热力图 dataplot = sns.heatmap(r, cmap="gist_heat", vmin=-0.2, vmax=1.1, annot=True, mask=mask, fmt='.4f', annot_kws={"fontsize":15}, cbar_kws={"shrink": 0.81}, clip_on=False, square=True) # 手动给下三角区域添加网格线 n = r.shape[0] for i in range(n): for j in range(i+1): # 仅遍历下三角单元格 # 获取单元格的坐标范围 x0, x1 = dataplot.get_xticks()[j], dataplot.get_xticks()[j+1] y0, y1 = dataplot.get_yticks()[n-1-i], dataplot.get_yticks()[n-i] # 绘制单元格的四条边框 plt.plot([x0, x0, x1, x1, x0], [y0, y1, y1, y0, y0], color='black', linewidth=1) plt.xticks(rotation=90) plt.yticks(rotation=0) plt.savefig('/content/heatmap.png', bbox_inches='tight', dpi=330) plt.show()
关键改动说明
- 移除了原代码中
heatmap的linewidths和linecolor参数,避免给所有单元格添加边框 - 通过遍历下三角区域的单元格坐标,手动绘制每个单元格的黑色边框
- 利用
dataplot.get_xticks()和dataplot.get_yticks()获取单元格的位置范围,确保边框位置准确
内容的提问来源于stack exchange,提问作者Abid Morshed
相关产品推荐
相关产品推荐

