如何为Seaborn山脊图添加均值、标准差线及阴影,移除y轴'Density'标签
问题与解决方案
需求与问题
- 数据包含多嵌套类别,调整可视化方案:放弃堆叠密度图的山脊组合,改为为每个变量绘制山脊图,要求:
- 增加均值线、标准差线
- 填充两条标准差线之间的曲线区域
- 基于企鹅数据集实现Seaborn山脊图后,无法移除y轴上的"Density"标签
原测试代码
# seaborn ridge plots with penguins dataset import logging import pandas as pd import matplotlib.pyplot as plt import numpy as np import seaborn as sns import os import errno LOG_FORMAT=("%(levelname) -5s time:%(asctime)s [%(funcName) ""-5s %(lineno) -5d]: %(message)s"); logging.basicConfig(level=logging.INFO, format=LOG_FORMAT); LOGGER = logging.getLogger(__name__); logger_obj: logging.Logger=LOGGER; my_df = sns.load_dataset("penguins"); sns.set_theme(style="white", rc={"axes.facecolor": (1, 1, 1, 1)});#background transparency def mkdir_p(path): if(not(os.path.exists(path) and os.path.isdir(path))): try: os.makedirs(path,exist_ok=True); except OSError as exc: # Python >2.5 if exc.errno == errno.EEXIST and os.path.isdir(path): pass; else: raise exc; def generate_plot( logger_obj: logging.Logger ,my_df: pd.DataFrame ,sample_size: int ,axs2 ): my_df2 = my_df.copy(deep=True); species_list: list=list(my_df2["species"].unique()); my_df3: pd.DataFrame; sample_size2: int=sample_size; for i2, species in enumerate(species_list): species_record_count=len(my_df2[my_df2["species"]==species]); flipper_length_mm_sum=my_df2[(my_df2["species"]==species)]["flipper_length_mm"].sum(); logger_obj.info("species is :'{0}', count is:{1}, flipper_length_mm_sum is:{2}".format(species, species_record_count, flipper_length_mm_sum)); if sample_size2>species_record_count: sample_size2=species_record_count; for i2, species in enumerate(species_list): my_df4=my_df2[my_df2["species"]==species].sample(sample_size2); species_record_count=len(my_df4); flipper_length_mm_sum=my_df4["flipper_length_mm"].sum(); logger_obj.info("species is :'{0}', count is:{1}, flipper_length_mm_sum is:{2}".format(species, species_record_count, flipper_length_mm_sum)); if i2==0: my_df3=my_df4[:]; else: my_df3=pd.concat([my_df3, my_df4], ignore_index=True); if 1==1: sns.set_theme(style="white", rc={"axes.facecolor": (0, 0, 0, 0), 'axes.linewidth':2}); palette = sns.color_palette("Set2", 12); g = sns.FacetGrid(data=my_df3, palette=palette, row="species", hue="species", aspect=9, height=1.2) sns.set_theme(style="white", rc={"axes.facecolor": (0, 0, 0, 0)}); g.map_dataframe(sns.kdeplot, x="flipper_length_mm", fill=True, alpha=1); g.map_dataframe(sns.kdeplot, x="flipper_length_mm", color="white"); def label_f(x, color, label): ax2=plt.gca(); ax2.text(0, .2, label, color="black", fontsize=13, ha="left", va="center", transform=ax2.transAxes); g.map(label_f, "species"); g.fig.subplots_adjust(hspace=-.5); g.set_titles(""); g.set(yticks=[], xlabel="flipper_length_mm"); g.set_titles(col_template="", row_template=""); g.despine(left=True); image_png_fn: str="images/penguins.ridge_plot/sample_day_feature.flipper_length_mm.all_species.png"; logger_obj.info("image_png_fn is :'{0}'".format(image_png_fn)); mkdir_p(os.path.abspath(os.path.join(image_png_fn, os.pardir))); plt.savefig(image_png_fn); image_png_fn=None; sample_size: int=30000; generate_plot( logger_obj ,my_df ,sample_size ,None );
修改后代码(解决标签问题+实现均值/标准差需求)
# seaborn ridge plots with penguins dataset import logging import pandas as pd import matplotlib.pyplot as plt import numpy as np import seaborn as sns import os import errno LOG_FORMAT=("%(levelname) -5s time:%(asctime)s [%(funcName) ""-5s %(lineno) -5d]: %(message)s"); logging.basicConfig(level=logging.INFO, format=LOG_FORMAT); LOGGER = logging.getLogger(__name__); logger_obj: logging.Logger=LOGGER; my_df = sns.load_dataset("penguins"); sns.set_theme(style="white", rc={"axes.facecolor": (1, 1, 1, 1)});#background transparency def mkdir_p(path): if(not(os.path.exists(path) and os.path.isdir(path))): try: os.makedirs(path,exist_ok=True); except OSError as exc: # Python >2.5 if exc.errno == errno.EEXIST and os.path.isdir(path): pass; else: raise exc; def generate_plot( logger_obj: logging.Logger ,my_df: pd.DataFrame ,sample_size: int ,axs2 ): my_df2 = my_df.copy(deep=True); species_list: list=list(my_df2["species"].unique()); my_df3: pd.DataFrame; sample_size2: int=sample_size; for i2, species in enumerate(species_list): species_record_count=len(my_df2[my_df2["species"]==species]); flipper_length_mm_sum=my_df2[(my_df2["species"]==species)]["flipper_length_mm"].sum(); logger_obj.info("species is :'{0}', count is:{1}, flipper_length_mm_sum is:{2}".format(species, species_record_count, flipper_length_mm_sum)); if sample_size2>species_record_count: sample_size2=species_record_count; for i2, species in enumerate(species_list): my_df4=my_df2[my_df2["species"]==species].sample(sample_size2); species_record_count=len(my_df4); flipper_length_mm_sum=my_df4["flipper_length_mm"].sum(); logger_obj.info("species is :'{0}', count is:{1}, flipper_length_mm_sum is:{2}".format(species, species_record_count, flipper_length_mm_sum)); if i2==0: my_df3=my_df4[:]; else: my_df3=pd.concat([my_df3, my_df4], ignore_index=True); sns.set_theme(style="white", rc={"axes.facecolor": (0, 0, 0, 0), 'axes.linewidth':2}); palette = sns.color_palette("Set2", 12); g = sns.FacetGrid(data=my_df3, palette=palette, row="species", hue="species", aspect=9, height=1.2) # 绘制KDE曲线 g.map_dataframe(sns.kdeplot, x="flipper_length_mm", fill=True, alpha=1); g.map_dataframe(sns.kdeplot, x="flipper_length_mm", color="white"); # 定义绘制均值、标准差及填充的函数 def add_stats(x, color, label): ax = plt.gca() mean_val = x.mean() std_val = x.std() # 绘制均值线 ax.axvline(mean_val, color="red", linestyle="--", linewidth=1.5) # 绘制±1标准差线 ax.axvline(mean_val - std_val, color="blue", linestyle=":", linewidth=1) ax.axvline(mean_val + std_val, color="blue", linestyle=":", linewidth=1) # 获取KDE曲线数据,填充标准差区间 kde = sns.kdeplot(x=x, ax=ax, color=color, alpha=0) x_kde = kde.get_lines()[0].get_xdata() y_kde = kde.get_lines()[0].get_ydata() # 筛选标准差区间内的曲线部分 mask = (x_kde >= mean_val - std_val) & (x_kde <= mean_val + std_val) ax.fill_between(x_kde[mask], y_kde[mask], color=color, alpha=0.3) # 应用统计线绘制 g.map(add_stats, "flipper_length_mm") def label_f(x, color, label): ax2=plt.gca(); ax2.text(0, .2, label, color="black", fontsize=13, ha="left", va="center", transform=ax2.transAxes); g.map(label_f, "species"); g.fig.subplots_adjust(hspace=-.5); g.set_titles(""); # 移除y轴标签和刻度 g.set(yticks=[], xlabel="flipper_length_mm", ylabel="") g.set_titles(col_template="", row_template=""); g.despine(left=True); # 遍历子轴确保y轴标签清空(兜底处理) for ax in g.axes.flat: ax.set_ylabel("") image_png_fn: str="images/penguins.ridge_plot/sample_day_feature.flipper_length_mm.all_species.png"; logger_obj.info("image_png_fn is :'{0}'".format(image_png_fn)); mkdir_p(os.path.abspath(os.path.join(image_png_fn, os.pardir))); plt.savefig(image_png_fn); image_png_fn=None; sample_size: int=30000; generate_plot( logger_obj ,my_df ,sample_size ,None );
关键修改说明
- 移除y轴"Density"标签:
- 在
g.set()中添加ylabel="",直接清空y轴标签 - 额外遍历所有子轴执行
ax.set_ylabel(""),避免个别子轴标签残留
- 在
- 实现均值/标准差需求:
- 新增
add_stats函数,计算当前分组的均值和标准差 - 绘制红色虚线作为均值线,蓝色点线作为±1标准差线
- 获取KDE曲线数据,填充标准差区间内的区域,增强视觉辨识度
- 新增
内容的提问来源于stack exchange,提问作者Allan K
相关产品推荐
相关产品推荐

