如何在组合柱状图与热力图时对齐Y轴刻度标签
解决Seaborn多子图(柱状图+热力图)Y轴对齐问题
问题描述
使用Seaborn组合两个横向柱状图与一个热力图时,开启sharey="row"后出现以下问题:
- Y轴标签与右侧热力图对齐
- 左侧柱状图顶部被截断,且柱状图与Y轴标签无法对齐
原始可运行代码如下:
import numpy as np import pandas as pd import seaborn as sns import matplotlib.pyplot as plt from matplotlib.colors import LogNorm ### Generate example data np.random.seed(123) year = [2018, 2019, 2020, 2021] task = [x + 2 for x in range(18)] student = [x for x in range(200)] amount = [x + 10 for x in range(90)] violation = [letter for letter in "thisisjustsampletextforlabels"] # one letter labels df_example = pd.DataFrame({ # some ways to create random data 'year':np.random.choice(year,500), 'task':np.random.choice(task,500), 'violation':np.random.choice(violation, 500), 'amount':np.random.choice(amount, 500), 'student':np.random.choice(student, 500) }) ### My code temp = df_example.groupby(["violation"])["amount"].sum().sort_values(ascending = False).reset_index() total_violations = temp["amount"].sum() sns.set(font_scale = 1.2) f, axs = plt.subplots(1,3, figsize=(5,5), sharey="row", gridspec_kw=dict(width_ratios=[3,1.5,5])) # Plot frequency df1 = df_example.groupby(["year","violation"])["amount"].sum().sort_values(ascending = False).reset_index() frequency = sns.barplot(data = df1, y = "violation", x = "amount", log = True, ax=axs[0]) # Plot percent df2 = df_example.groupby(["violation"])["amount"].sum().sort_values(ascending = False).reset_index() total_violations = df2["amount"].sum() percent = sns.barplot(x='amount', y='violation', estimator=lambda x: sum(x) / total_violations * 100, data=df2, ax=axs[1]) # Pivot table and plot heatmap df_heatmap = df_example.groupby(["violation", "task"])["amount"].sum().sort_values(ascending = False).reset_index() df_heatmap_pivot = df_heatmap.pivot("violation", "task", "amount") df_heatmap_pivot = df_heatmap_pivot.reindex(index=df_heatmap["violation"].unique()) heatmap = sns.heatmap(df_heatmap_pivot, fmt = "d", cmap="Greys", norm=LogNorm(), ax=axs[2]) plt.subplots_adjust(top=1) axs[2].set_facecolor('xkcd:white') axs[2].set(ylabel="",xlabel="Task") axs[0].set_xlabel('Total amount of violations per year') axs[1].set_xlabel('Percent (%)') axs[1].set_ylabel('') axs[0].set_ylabel('Violation')
解决方案
核心问题是Seaborn热力图会自动调整轴的边距和范围,破坏了sharey的对齐效果。通过以下修改解决:
修改后的完整代码:
import numpy as np import pandas as pd import seaborn as sns import matplotlib.pyplot as plt from matplotlib.colors import LogNorm ### Generate example data np.random.seed(123) year = [2018, 2019, 2020, 2021] task = [x + 2 for x in range(18)] student = [x for x in range(200)] amount = [x + 10 for x in range(90)] violation = [letter for letter in "thisisjustsampletextforlabels"] # one letter labels df_example = pd.DataFrame({ # some ways to create random data 'year':np.random.choice(year,500), 'task':np.random.choice(task,500), 'violation':np.random.choice(violation, 500), 'amount':np.random.choice(amount, 500), 'student':np.random.choice(student, 500) }) ### Modified code temp = df_example.groupby(["violation"])["amount"].sum().sort_values(ascending = False).reset_index() total_violations = temp["amount"].sum() sns.set(font_scale = 1.2) # 增大figsize避免截断,优化宽度比例 f, axs = plt.subplots(1,3, figsize=(12,8), sharey="row", gridspec_kw=dict(width_ratios=[3,1.5,5])) # Plot frequency df1 = df_example.groupby(["year","violation"])["amount"].sum().sort_values(ascending = False).reset_index() frequency = sns.barplot(data = df1, y = "violation", x = "amount", log = True, ax=axs[0]) # Plot percent df2 = df_example.groupby(["violation"])["amount"].sum().sort_values(ascending = False).reset_index() total_violations = df2["amount"].sum() percent = sns.barplot(x='amount', y='violation', estimator=lambda x: sum(x) / total_violations * 100, data=df2, ax=axs[1]) # Pivot table and plot heatmap df_heatmap = df_example.groupby(["violation", "task"])["amount"].sum().sort_values(ascending = False).reset_index() df_heatmap_pivot = df_heatmap.pivot("violation", "task", "amount") df_heatmap_pivot = df_heatmap_pivot.reindex(index=df_heatmap["violation"].unique()) # 绘制热力图时禁用颜色条(可选,如需颜色条需手动调整位置),并设置yticklabels匹配 heatmap = sns.heatmap(df_heatmap_pivot, fmt = "d", cmap="Greys", norm=LogNorm(), ax=axs[2], cbar=False, yticklabels=True) # 同步所有子图的Y轴范围,确保对齐 y_min, y_max = axs[0].get_ylim() axs[1].set_ylim(y_min, y_max) axs[2].set_ylim(y_min, y_max) # 调整子图间距,避免重叠 plt.subplots_adjust(top=0.95, bottom=0.05, wspace=0.2) axs[2].set_facecolor('xkcd:white') axs[2].set(ylabel="",xlabel="Task") axs[0].set_xlabel('Total amount of violations per year') axs[1].set_xlabel('Percent (%)') axs[1].set_ylabel('') axs[0].set_ylabel('Violation') plt.show()
关键修改说明
- 增大
figsize为(12,8),避免顶部内容被截断; - 热力图添加
cbar=False(若需要颜色条,可使用plt.colorbar(heatmap.collections[0], ax=axs[2])手动放置); - 获取第一个柱状图的Y轴范围,强制所有子图使用相同范围,确保对齐;
- 调整
plt.subplots_adjust()的参数,优化子图间距。
内容的提问来源于stack exchange,提问作者Linus Östlund
相关产品推荐
相关产品推荐

