线性回归置信区间绘图:如何在Seaborn风格图中添加统计信息至图例
解决方案:Matplotlib绘制带置信区间的线性回归图(附统计信息)
方法一:Matplotlib + Statsmodels(手动实现置信区间,完整获取统计数据)
Statsmodels可直接输出详细回归统计结果,同时支持计算置信区间,完全匹配需求:
- 拟合模型并提取统计数据
import numpy as np import matplotlib.pyplot as plt import statsmodels.api as sm # 生成示例数据 x = np.random.rand(100) y = 2 * x + 1 + np.random.randn(100)*0.2 # 添加截距项(Statsmodels需手动指定) x_with_const = sm.add_constant(x) model = sm.OLS(y, x_with_const).fit() # 提取核心统计数据:R²、斜率、截距、p值等 print(model.summary()) # 完整统计报告 r_squared = model.rsquared slope = model.params[1] intercept = model.params[0] p_value = model.pvalues[1]
- 计算置信区间并绘图(复刻Seaborn风格)
# 生成绘图用的连续x序列 x_plot = np.linspace(x.min(), x.max(), 100) x_plot_with_const = sm.add_constant(x_plot) # 预测拟合值与95%置信区间 y_pred = model.predict(x_plot_with_const) pred_ci = model.get_prediction(x_plot_with_const).conf_int() # 启用Seaborn风格并绘图 plt.style.use('seaborn-v0_8') plt.scatter(x, y, alpha=0.6, label='原始数据') plt.plot(x_plot, y_pred, 'r-', label=f'回归线 (y={slope:.2f}x+{intercept:.2f}, R²={r_squared:.2f})') plt.fill_between(x_plot, pred_ci[:,0], pred_ci[:,1], color='r', alpha=0.2, label='95%置信区间') # 完善图表标注 plt.xlabel('X') plt.ylabel('Y') plt.legend() plt.show()
该方案完全手动控制绘图细节,统计数据可直接从模型对象提取,置信区间通过fill_between实现,效果与Seaborn的regplot一致。
方法二:pydove.regplot 添加统计信息到图例
pydove的regplot已内置回归线与置信区间绘制,只需额外提取统计数据并手动添加图例条目即可:
import numpy as np import matplotlib.pyplot as plt from pydove import regplot from sklearn.linear_model import LinearRegression from sklearn.metrics import r2_score # 示例数据 x = np.random.rand(100) y = 2 * x + 1 + np.random.randn(100)*0.2 # 用pydove绘制回归图 ax = regplot(x, y) # 拟合模型获取统计数据 model = LinearRegression().fit(x.reshape(-1,1), y) r_squared = r2_score(y, model.predict(x.reshape(-1,1))) slope = model.coef_[0] intercept = model.intercept_ # 创建自定义图例条目 from matplotlib.lines import Line2D custom_lines = [ Line2D([0], [0], color='C0', lw=2), # 匹配回归线颜色 Line2D([0], [0], color='C0', marker='o', lw=0, alpha=0.6), # 匹配散点样式 Line2D([0], [0], color='C0', lw=0, alpha=0.2) # 匹配置信区间填充 ] labels = [ f'回归线 (y={slope:.2f}x+{intercept:.2f})', '原始数据', f'95%置信区间 (R²={r_squared:.2f})' ] # 更新图例 ax.legend(custom_lines, labels) plt.xlabel('X') plt.ylabel('Y') plt.show()
若不想重复拟合模型,也可直接从pydove绘图后返回的Axes对象中提取线条属性,手动拼接统计信息到图例文本中。
内容的提问来源于stack exchange,提问作者MateaMar
相关产品推荐
相关产品推荐

