如何使用Seaborn绘制两个分布的差值图?
绘制两个KDE分布的差值(disA - disB)
要实现你的需求,核心思路是先分别获取两个分布的概率密度估计值,再计算它们的差值,最后用matplotlib绘制差值曲线。Seaborn的kdeplot默认直接绘图而不返回密度数据,所以我们可以用scipy.stats.gaussian_kde来更灵活地处理数值计算,具体步骤如下:
步骤1:准备数据并生成KDE模型
首先提取两组数据,创建KDE对象来估计密度:
import numpy as np import scipy.stats as stats import matplotlib.pyplot as plt import seaborn as sns # 提取两组数据,注意处理缺失值 data_outcome0 = df['term'][df['outcome'] == 0].dropna().values data_outcome1 = df['term'][df['outcome'] == 1].dropna().values # 创建KDE密度估计对象 kde_outcome0 = stats.gaussian_kde(data_outcome0) kde_outcome1 = stats.gaussian_kde(data_outcome1)
步骤2:生成统一的X轴并计算差值
为了确保两个密度在相同的X点上计算,我们生成覆盖两组数据范围的X轴,然后计算密度差值:
# 确定X轴的范围,覆盖两组数据的最小和最大值 x_min = min(data_outcome0.min(), data_outcome1.min()) x_max = max(data_outcome0.max(), data_outcome1.max()) x = np.linspace(x_min, x_max, 1000) # 生成1000个均匀分布的X点 # 计算两组数据在X轴上的密度值 density0 = kde_outcome0(x) density1 = kde_outcome1(x) # 计算差值:disA (outcome=0) - disB (outcome=1) density_diff = density0 - density1
步骤3:绘制差值曲线并可视化
最后绘制差值曲线,还可以填充正负区域让结果更直观:
plt.figure(figsize=(10, 6)) # 绘制差值曲线 plt.plot(x, density_diff, color='darkblue', linewidth=2, label='Density Difference (0 - 1)') # 填充差值为正的区域(disA > disB) plt.fill_between(x, density_diff, 0, where=density_diff > 0, color='red', alpha=0.3) # 填充差值为负的区域(disA < disB) plt.fill_between(x, density_diff, 0, where=density_diff < 0, color='green', alpha=0.3) # 添加标签和标题 plt.xlabel('term') plt.ylabel('Density Difference') plt.title('Difference Between KDE Distributions (outcome=0 minus outcome=1)') plt.legend() plt.grid(alpha=0.3) plt.show()
补充说明
如果你坚持想用Seaborn的kdeplot来获取数据,可以先绘制两个KDE曲线,然后从Axes对象的lines属性中提取X和Y数据,但这种方法依赖于Seaborn的内部实现,不如scipy的方法稳定:
fig, ax = plt.subplots() sns.kdeplot(df['term'][df['outcome'] == 0], shade=False, color='red', ax=ax) sns.kdeplot(df['term'][df['outcome'] == 1], shade=False, color='green', ax=ax) # 提取两条曲线的数据 line0 = ax.lines[0] line1 = ax.lines[1] x0, y0 = line0.get_xdata(), line0.get_ydata() x1, y1 = line1.get_xdata(), line1.get_ydata() # 因为X轴范围一致,可以直接计算差值(如果不一致需要插值对齐) diff = y0 - y1 # 绘制差值曲线 plt.figure() plt.plot(x0, diff) plt.show()
内容的提问来源于stack exchange,提问作者mllamazares
相关产品推荐
相关产品推荐

