如何在seaborn热力图Y轴右侧为多行分组添加标签或图例
问题描述
我用Seaborn绘制了按行分组、每组对应不同颜色的热力图,想在Y轴右侧为每个颜色组添加标签(比如红色组标注“上升趋势”,紫色组标注“新兴趋势”),同时为各颜色组添加图例,求实现方法。
原代码如下:
import pandas as pd import seaborn as sns import matplotlib.pyplot as plt agg_df=pd.DataFrame(agg, columns=a_list, index=agg_x) print(agg_df) agg_df_1=agg_df.copy() agg_df_2=agg_df.copy() agg_df_3=agg_df.copy() agg_df_4=agg_df.copy() #generate heatmap agg_df_1.iloc[:,24:] = float('nan') g=sns.heatmap(agg_df_1.transpose(), annot=False,fmt=".1f",annot_kws={"fontsize":6},cmap=plt.cm.Reds,xticklabels=1, yticklabels=1,cbar=False) g.set_xticklabels(g.get_xticklabels(), fontsize = 8) g.set_yticklabels(g.get_yticklabels(), fontsize = 8) agg_df_2.iloc[:,:24]=float('nan') agg_df_2.iloc[:,35:]=float('nan') g=sns.heatmap(agg_df_2.transpose(), annot=False,fmt=".1f",annot_kws={"fontsize":6},cmap=plt.cm.Oranges,xticklabels=1, yticklabels=1,cbar=False) g.set_xticklabels(g.get_xticklabels(), fontsize = 8) g.set_yticklabels(g.get_yticklabels(), fontsize = 8) agg_df_3.iloc[:,:35]=float('nan') agg_df_3.iloc[:,44:]=float('nan') g=sns.heatmap(agg_df_3.transpose(), annot=False,fmt=".1f",annot_kws={"fontsize":6},cmap=plt.cm.Purples,xticklabels=1, yticklabels=1,cbar=False) g.set_xticklabels(g.get_xticklabels(), fontsize = 8) g.set_yticklabels(g.get_yticklabels(), fontsize = 8) agg_df_4.iloc[:,:44]=float('nan') g=sns.heatmap(agg_df_4.transpose(), annot=False,fmt=".1f",annot_kws={"fontsize":6},cmap=plt.cm.Greens,xticklabels=1, yticklabels=1,cbar=False) g.set_xticklabels(g.get_xticklabels(), fontsize = 8) g.set_yticklabels(g.get_yticklabels(), fontsize = 8) for label in g.get_yticklabels(): label.set_weight('bold') for label in g.get_xticklabels(): label.set_weight('bold') #g.set(ylabel='Decreasing Emerging Mix Increasing') plt.show()
实现方法
1. Y轴右侧添加分组标签
利用Matplotlib的text()函数,在热力图轴的右侧对应每组行的中间位置添加文本标签。先确定每组行的范围,计算每组的垂直中心坐标,再设置文本的位置和样式。
2. 添加颜色组图例
通过创建自定义的颜色补丁(matplotlib.patches.Patch),每个补丁对应一个颜色组的主色调,再用plt.legend()添加图例。
修改后的完整代码
import pandas as pd import seaborn as sns import matplotlib.pyplot as plt from matplotlib.patches import Patch # 假设agg、a_list、agg_x已定义 agg_df = pd.DataFrame(agg, columns=a_list, index=agg_x) # 定义分组信息:(列范围, 颜色映射, 标签) groups = [ (slice(None, 24), plt.cm.Reds, "上升趋势"), (slice(24, 35), plt.cm.Oranges, "新兴趋势"), (slice(35, 44), plt.cm.Purples, "混合趋势"), (slice(44, None), plt.cm.Greens, "下降趋势") ] # 初始化画布和轴 fig, ax = plt.subplots(figsize=(12, 8)) # 循环绘制每个分组的热力图 for col_slice, cmap, label in groups: temp_df = agg_df.copy() # 只保留当前分组的列,其余设为NaN temp_df.loc[:, ~temp_df.columns.isin(temp_df.columns[col_slice])] = float('nan') sns.heatmap(temp_df.transpose(), annot=False, fmt=".1f", annot_kws={"fontsize":6}, cmap=cmap, xticklabels=1, yticklabels=1, cbar=False, ax=ax) # 设置坐标轴标签样式 ax.set_xticklabels(ax.get_xticklabels(), fontsize=8, weight='bold') ax.set_yticklabels(ax.get_yticklabels(), fontsize=8, weight='bold') # 添加Y轴右侧的分组标签 y_ticks = ax.get_yticks() total_rows = len(y_ticks) - 1 # 刻度数比行数多1 # 计算每组行的起始和结束索引(transpose后行对应原列) group_ranges = [ (0, 24), (24, 35), (35, 44), (44, total_rows) ] group_labels = ["上升趋势", "新兴趋势", "混合趋势", "下降趋势"] for (start, end), label in zip(group_ranges, group_labels): # 计算组的垂直中心位置 mid_y = (y_ticks[start] + y_ticks[end]) / 2 # 在轴右侧添加文本,x=1.02是相对轴的位置(1是轴右端) ax.text(1.02, mid_y, label, va='center', ha='left', weight='bold', transform=ax.get_yaxis_transform()) # 添加自定义图例 legend_elements = [Patch(facecolor=cmap(0.7), edgecolor='black', label=label) for _, cmap, label in groups] ax.legend(handles=legend_elements, bbox_to_anchor=(1.15, 1), loc='upper left') plt.tight_layout() plt.show()
说明
- 把重复的绘图逻辑改成循环,减少代码冗余;
- 分组标签通过
ax.text()添加,transform=ax.get_yaxis_transform()保证文本位置随Y轴缩放而不变; - 图例用自定义Patch创建,
bbox_to_anchor参数控制图例位置,避免遮挡热力图; - 可根据实际分组范围和标签内容调整
groups列表中的参数。
内容的提问来源于stack exchange,提问作者Travelling Salesman
相关产品推荐
相关产品推荐

