如何在带hue的Seaborn柱状图内显示观测数量(n)
在Seaborn柱状图柱子内显示分组观测数量(n)
问题背景
需要创建基于%_time指标的带hue分组柱状图,同时在每个柱子内部显示对应分组的观测数量(n)。现有DataFrame如下:
| Treatment | Condition | %_time |
|---|---|---|
| STZ | Stressed | 3 |
| Control | Stressed | 6 |
| STZ | Unstressed | 2 |
| Control | Unstressed | 8 |
现有实现代码
已成功生成目标柱状图,但还未添加n值标签,代码如下:
import seaborn as sns import matplotlib.pyplot as plt color = sns.color_palette("Paired") sns.set_style(style='white') fig, ax = plt.subplots(figsize=(9,6)) ax = sns.barplot(x='Treatment', y='%_time', hue='Condition', data=df, capsize=.1, palette=color) plt.legend(title='Groups', loc='upper right') plt.xlabel("Treatment") plt.ylabel("% Time in Open Arm") plt.title("Stress in STZ vs Vehicle ", size=14)
遇到的问题
- 使用
countplot能显示n值,但y轴会变为计数,无法保留%_time的数值展示; - 直接给
barplot添加bar_label,显示的是柱子对应的%_time数值,而非分组的观测数n。
解决方案
1. 先统计每个分组的观测数
先按Treatment和Condition组合分组,计算每组的样本量:
# 统计每个(Treatment, Condition)组合的观测数 n_counts = df.groupby(['Treatment', 'Condition']).size().reset_index(name='n')
2. 遍历柱子添加n值标签
在原有绘图代码基础上,遍历每个柱子,匹配对应的分组信息,将n值添加到柱子内部:
import seaborn as sns import matplotlib.pyplot as plt color = sns.color_palette("Paired") sns.set_style(style='white') fig, ax = plt.subplots(figsize=(9,6)) # 绘制基础柱状图 ax = sns.barplot(x='Treatment', y='%_time', hue='Condition', data=df, capsize=.1, palette=color) # 遍历所有柱子,添加n值标签 for p in ax.patches: # 计算柱子中心的坐标位置 x_center = p.get_x() + p.get_width() / 2 y_center = p.get_height() / 2 # 标签放在柱子垂直中心,可根据需求调整 # 根据柱子位置判断所属的Treatment和Condition分组 treatment_idx = int(p.get_x() // (p.get_width() * 2)) condition_idx = int((p.get_x() % (p.get_width() * 2)) // p.get_width()) # 获取对应的分组名称 treatment = df['Treatment'].unique()[treatment_idx] condition = df['Condition'].unique()[condition_idx] # 匹配对应的n值 n_value = n_counts[(n_counts['Treatment'] == treatment) & (n_counts['Condition'] == condition)]['n'].values[0] # 添加文本标签 ax.text(x_center, y_center, f'n={n_value}', ha='center', va='center', fontweight='bold') # 图表美化设置 plt.legend(title='Groups', loc='upper right') plt.xlabel("Treatment") plt.ylabel("% Time in Open Arm") plt.title("Stress in STZ vs Vehicle ", size=14) plt.show()
注意事项
- 若原始DataFrame中每个
(Treatment, Condition)组包含多条数据,groupby.size()会自动统计每组的真实观测数,无需修改代码; - 标签位置可通过调整
y_center的计算方式修改,比如想让标签靠近柱子顶部,可改为y_center = p.get_height() * 0.8。
内容的提问来源于stack exchange,提问作者Adriana
相关产品推荐
相关产品推荐

