You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何仅在热力图的下三角区域显示网格?

问题:热力图仅在下三角区域保留网格

我制作了一个遮罩上三角部分的热力图,但设置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()

关键改动说明

  1. 移除了原代码中heatmap的linewidths和linecolor参数,避免给所有单元格添加边框
  2. 通过遍历下三角区域的单元格坐标,手动绘制每个单元格的黑色边框
  3. 利用dataplot.get_xticks()和dataplot.get_yticks()获取单元格的位置范围,确保边框位置准确

内容的提问来源于stack exchange,提问作者Abid Morshed

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.17 16:12:47