You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何为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
);

关键修改说明

  1. 移除y轴"Density"标签:
    • 在g.set()中添加ylabel="",直接清空y轴标签
    • 额外遍历所有子轴执行ax.set_ylabel(""),避免个别子轴标签残留
  2. 实现均值/标准差需求:
    • 新增add_stats函数,计算当前分组的均值和标准差
    • 绘制红色虚线作为均值线,蓝色点线作为±1标准差线
    • 获取KDE曲线数据,填充标准差区间内的区域,增强视觉辨识度

内容的提问来源于stack exchange,提问作者Allan K

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.18 20:00:53