如何修改Seaborn代码实现正确分面堆叠柱状图(解决>7分面异常)
问题:Seaborn FacetGrid堆叠柱状图异常修复
数据与需求
现有如下Pandas DataFrame数据:
d = {'cutoff': ['>6']*8+ ['>7']*8 +['>8']*8, 'month': ['3m','3m','3m','3m','6m','6m','6m','6m'] * 3, 'outcome': ['D','D','I','I']*6, 'count': ['0','>0']*12, 'proportion': [0.25, 0.75, 0.4, 0.6, 0.3, 0.7, 0.8, 0.2, 0.1, 0.9, 0, 1.0, 0.15, 0.85, 0.2, 0.8, 0.3, 0.7, 0.8, 0.2, 0.35, 0.65, 1.0, 0]} tdf = pd.DataFrame(d)
需求是展示不同cutoff、不同month下,不同outcome对应的count占比差异,需创建FacetGrid分面堆叠柱状图,预期每个cutoff分面下,按month+outcome分组展示堆叠的占比柱形。
问题代码与异常
使用以下代码绘制时,>7分面的柱状图分布异常(呈均匀分布):
def draw_stacked_bar(*args, **kwargs): data = kwargs.pop('data') ax = plt.gca() pivot_df = data.pivot(index=['month','outcome'], columns='gitis_count', values='proportion') pivot_df.plot(kind='bar', stacked=True, ax=ax) plt.xticks(rotation=45) g = sb.FacetGrid(grouped_df, col='cutoff', col_wrap=3, height=4, sharey=True) g.map_dataframe(draw_stacked_bar)
问题分析与修复代码
问题原因
- 代码中使用未定义的
grouped_df,应替换为原始数据tdf; - 透视时列名写错:
gitis_count应为数据中的count; - 不同分面的透视列顺序、分组索引不一致,导致堆叠逻辑混乱,部分分组缺失引发柱形错位。
修正后的代码
import seaborn as sb import pandas as pd import matplotlib.pyplot as plt d = {'cutoff': ['>6']*8+ ['>7']*8 +['>8']*8, 'month': ['3m','3m','3m','3m','6m','6m','6m','6m'] * 3, 'outcome': ['D','D','I','I']*6, 'count': ['0','>0']*12, 'proportion': [0.25, 0.75, 0.4, 0.6, 0.3, 0.7, 0.8, 0.2, 0.1, 0.9, 0, 1.0, 0.15, 0.85, 0.2, 0.8, 0.3, 0.7, 0.8, 0.2, 0.35, 0.65, 1.0, 0]} tdf = pd.DataFrame(d) # 固定count的堆叠顺序,确保所有分面一致 count_order = ['0', '>0'] # 固定分组索引顺序,避免缺失类别导致错位 index_order = [('3m', 'D'), ('3m', 'I'), ('6m', 'D'), ('6m', 'I')] def draw_stacked_bar(*args, **kwargs): data = kwargs.pop('data') ax = plt.gca() # 透视后重新索引,统一所有分面的分组和列顺序 pivot_df = data.pivot(index=['month','outcome'], columns='count', values='proportion') pivot_df = pivot_df.reindex(index=index_order, columns=count_order).fillna(0) pivot_df.plot(kind='bar', stacked=True, ax=ax) # 自定义x轴标签,更贴合需求 ax.set_xticklabels([f"{m}-{o}" for m, o in index_order], rotation=45) ax.set_xlabel('') ax.legend(title='Count') g = sb.FacetGrid(tdf, col='cutoff', col_wrap=3, height=4, sharey=True) g.map_dataframe(draw_stacked_bar) plt.tight_layout() plt.show()
修复说明
- 修正了变量名和列名错误,确保数据来源正确;
- 强制指定
count堆叠顺序和分组索引顺序,避免分面间的展示逻辑不一致; - 用
fillna(0)处理缺失值,保证所有分组都有对应的占比数据; - 自定义x轴标签,让展示更符合预期需求。
内容的提问来源于stack exchange,提问作者pill45
相关产品推荐
相关产品推荐

