如何避免遍历multiple模式时重复执行sns.kdeplot计算?
解决方案:预计算KDE分布再复用
要解决大型数据集下重复计算的问题,核心是一次性算出每个类别的KDE分布数据,之后针对不同multiple模式直接复用这些数据绘图,避免多次调用kdeplot带来的冗余计算。
具体实现步骤
- 提取每个类别的数据子集,统一计算KDE曲线的x轴取值和对应密度值
- 针对
layer/stack/fill三种模式,分别对密度值做适配处理 - 用Matplotlib手动绘制不同样式的图表,全程复用预计算好的KDE数据
修改后的代码
import pandas as pd import numpy as np import matplotlib.pyplot as plt import seaborn as sns from scipy.stats import gaussian_kde np.random.seed(42) df = pd.DataFrame({ 'category': np.random.randint(0, 4, size=500), 'value': np.random.uniform(-10, 10, size=500) }) def multiplehue_kdeplots(df, val_col, cat_col, mults=['layer', 'stack', 'fill'], leg_idx=None, pal=None, **kwargs): # 1. 预计算每个类别的KDE数据 categories = sorted(df[cat_col].unique()) pal = pal or sns.color_palette(n_colors=len(categories)) # 统一x轴范围,保证所有KDE曲线用相同的x取值 all_values = df[val_col].values x_min, x_max = all_values.min() - 2, all_values.max() + 2 x = np.linspace(x_min, x_max, 1000) kde_data = [] for cat in categories: cat_values = df[df[cat_col] == cat][val_col].values kde = gaussian_kde(cat_values) y = kde(x) kde_data.append({ 'category': cat, 'x': x, 'y': y, 'color': pal[categories.index(cat)] }) # 2. 创建子图 fig, axs = plt.subplots(nrows=len(mults), sharex=True, **kwargs) axs = axs if len(mults) > 1 else [axs] # 3. 按不同模式绘制图表 for i, mode in enumerate(mults): ax = axs[i] current_base = np.zeros_like(x) for data in kde_data: if mode == 'layer': # 直接绘制单类别曲线与填充 ax.plot(data['x'], data['y'], color=data['color']) ax.fill_between(data['x'], 0, data['y'], color=data['color'], alpha=0.5) elif mode == 'stack': # 累加密度实现堆叠效果 ax.plot(data['x'], current_base + data['y'], color=data['color']) ax.fill_between(data['x'], current_base, current_base + data['y'], color=data['color'], alpha=0.5) current_base += data['y'] elif mode == 'fill': # 填充堆叠区域(无顶部曲线) ax.fill_between(data['x'], current_base, current_base + data['y'], color=data['color'], alpha=0.5) current_base += data['y'] ax.set_title(mode.title()) ax.set_ylabel('Density') # 仅在指定子图添加图例 if leg_idx == i: handles = [plt.Rectangle((0,0),1,1, color=d['color'], alpha=0.5) for d in kde_data] ax.legend(handles, [d['category'] for d in kde_data], title=cat_col) fig.subplots_adjust(hspace=0.3) plt.show() multiplehue_kdeplots(df, 'value', 'category')
关键说明
- 性能优化:大型数据集下仅需一次KDE计算,相比原函数的三次计算,时间开销直接降低约2/3
- 模式适配逻辑:
layer:独立绘制每个类别的密度曲线,互不叠加stack:累加每个类别的密度值,实现堆叠式的分层效果fill:与stack逻辑一致,仅保留填充区域(如需顶部曲线可自行添加)
- 一致性保障:统一x轴取值范围,确保不同模式下的图表对齐一致
内容的提问来源于stack exchange,提问作者bchate
相关产品推荐
相关产品推荐

