如何用Pandas为散点图添加线性回归线并查看线性相关性?
添加回归线与计算线性相关性解决方案
1. 依赖库准备
除pandas外,需导入matplotlib.pyplot用于绘图,scipy.stats计算回归参数;也可使用seaborn简化回归线绘制流程。
2. 手动计算回归并绘图
通过scipy.stats.linregress获取斜率、截距、相关系数等参数,在散点图基础上叠加回归线,同时输出变量间线性相关性指标:
import pandas as pd import matplotlib.pyplot as plt from scipy.stats import linregress # 加载数据集 df = pd.read_csv('/home/beki/Desktop/Final_report9.csv') # 定义变量配置:(y列名, 颜色, 标记) x_col = 'log_AGN_LUM(type-2)' y_configs = [ ('log_LUM_Ha', 'b', '^'), ('log_Lum_OII', 'r', 's'), ('log_LUM_H_beta', 'c', 'x'), ('log_LUM_OIII', 'g', 'o') ] # 创建绘图轴 fig, ax = plt.subplots(figsize=(10, 7)) # 遍历变量绘制散点+回归线,输出相关性 for y_col, color, marker in y_configs: # 过滤缺失值,避免回归计算报错 valid_data = df[[x_col, y_col]].dropna() x = valid_data[x_col] y = valid_data[y_col] # 绘制散点图 ax.scatter(x, y, color=color, marker=marker, alpha=0.5, label=y_col) # 计算线性回归参数 slope, intercept, r_value, p_value, std_err = linregress(x, y) # 生成回归线的x范围与对应y值 x_reg = [x.min(), x.max()] y_reg = [slope * xi + intercept for xi in x_reg] # 绘制回归线 ax.plot(x_reg, y_reg, color=color, linestyle='--', linewidth=1.5) # 打印相关性结果 print(f"{y_col} 与 {x_col} 的线性相关性:") print(f" 相关系数(r): {r_value:.4f}") print(f" p值: {p_value:.4f}\n") # 设置图表属性 ax.set_xlabel('AGN_LUM') ax.set_ylabel('Luminosity of emission lines') ax.set_title('AGN Luminosity Vs Luminosity of emission lines', weight='bold', size=12) ax.legend(loc="lower right") plt.show()
3. 快速查看全局相关性矩阵
用pandas.corr()生成所有目标变量的相关性矩阵,直观对比线性相关程度:
# 提取目标变量列并计算相关性矩阵 target_cols = [x_col] + [cfg[0] for cfg in y_configs] corr_matrix = df[target_cols].corr() print("变量相关性矩阵:") print(corr_matrix)
4. 简化方案:用Seaborn一键绘制
若无需手动控制回归计算细节,seaborn.regplot可直接生成带回归线的散点图:
import seaborn as sns import pandas as pd import matplotlib.pyplot as plt df = pd.read_csv('/home/beki/Desktop/Final_report9.csv') x_col = 'log_AGN_LUM(type-2)' y_cols = ['log_LUM_Ha', 'log_Lum_OII', 'log_LUM_H_beta', 'log_LUM_OIII'] colors = ['b', 'r', 'c', 'g'] markers = ['^', 's', 'x', 'o'] fig, ax = plt.subplots(figsize=(10, 7)) # 批量绘制带回归线的散点图 for y_col, color, marker in zip(y_cols, colors, markers): sns.regplot(x=x_col, y=y_col, data=df, ax=ax, scatter_kws={'color': color, 'marker': marker, 'alpha':0.5}, line_kws={'color': color, 'linestyle': '--'}, label=y_col) # 设置图表属性 ax.set_xlabel('AGN_LUM') ax.set_ylabel('Luminosity of emission lines') ax.set_title('AGN Luminosity Vs Luminosity of emission lines', weight='bold', size=12) ax.legend(loc="lower right") plt.show()
内容的提问来源于stack exchange,提问作者BEREKET ASSEFA
相关产品推荐
相关产品推荐

