使用subplot2grid绘制heatmap时如何保持子图列宽一致且去除多余标签
解决方案
1. 移除多余的"dates"标签
该标签是你的pivot数据表的列索引名称,有两种处理方式:
- 提前清除列索引名:执行
pivot.columns.name = None - 绘制后隐藏x轴标签:执行
ax1.set_xlabel('')
2. 子图水平对齐、共享x轴
子图错位的核心原因是你给两个热力图都设置了square=True,该参数会强制每个单元格为正方形,下方子图仅1行高度远小于上方7行的热力图,会自动压缩单元格宽度导致错位,搭配共享x轴的设置即可解决问题,完整修改代码如下:
import matplotlib.pyplot as plt import seaborn as sns import pandas as pd # 清除列索引名,直接移除dates标签 pivot.columns.name = None fig = plt.figure(figsize=(25, 15)) ax1 = plt.subplot2grid((23,20), (0,0), colspan=19, rowspan=17) # 共享ax1的x轴 ax2 = plt.subplot2grid((23,20), (19,0), colspan=19, rowspan=1, sharex=ax1) sns.set(font_scale=0.95) # 移除square=True参数,避免单元格强制正方形导致宽度错位 sns.heatmap(pivot, ax= ax1, annot=True, fmt=".0f", robust=True, linewidth=0.1, cmap="Blues") sns.heatmap((pd.DataFrame(pivot.sum(axis=0))).transpose(), ax=ax2, annot=True, fmt=".0f", robust=True, linewidth=0.1, cmap="Blues", xticklabels=False, yticklabels=False) # 隐藏上方热力图的x轴刻度,避免和下方子图重复 plt.setp(ax1.get_xticklabels(), visible=False) plt.show()
如果需要保留正方形单元格样式,可改用GridSpec设置匹配行数的高度比例,代码示例如下:
import matplotlib.pyplot as plt from matplotlib.gridspec import GridSpec import seaborn as sns import pandas as pd pivot.columns.name = None fig = plt.figure(figsize=(25, 15)) # 高度比例设置为7:1,匹配上下热力图的行数比,保证正方形单元格宽度一致 gs = GridSpec(2, 1, height_ratios=[7, 1], hspace=0.05) ax1 = fig.add_subplot(gs[0]) ax2 = fig.add_subplot(gs[1], sharex=ax1) sns.set(font_scale=0.95) # 可保留square=True参数 sns.heatmap(pivot, ax= ax1, annot=True, fmt=".0f", robust=True, linewidth=0.1, square=True, cmap="Blues") sns.heatmap((pd.DataFrame(pivot.sum(axis=0))).transpose(), ax=ax2, annot=True, fmt=".0f", robust=True, linewidth=0.1, square=True, cmap="Blues", xticklabels=False, yticklabels=False) plt.setp(ax1.get_xticklabels(), visible=False) plt.show()
内容的提问来源于stack exchange,提问作者bellotto
相关产品推荐
相关产品推荐

