Matplotlib线性回归拟合异常原因排查
拟合异常的根本原因分析
你的代码里导致回归线完全偏离的核心问题是对np.polyfit的输入处理错误:
问题核心代码
m, b = np.polyfit(X[0], y[0], 1)
你通过np.c_把X和y转换成了二维数组(形状为(n,1),n是样本总数),但X[0]只取了二维数组的第一行——也就是单个GDP数值(俄罗斯的9054.914),y[0]同理是单个生活满意度值(6)。用两个单点拟合直线,本质上是基于极少量数据生成了完全不贴合整体趋势的直线,再用全量X数组绘制这条错误直线,最终出现拟合严重偏差的结果。
正确处理方式
np.polyfit要求输入的X和y是一维数组,你可以通过两种方式修正:
- 把二维数组展平:
m, b = np.polyfit(X.flatten(), y.flatten(), 1)
- 创建
X和y时直接生成一维数组,不使用np.c_:
X = country_stats["GDP per capita"].values y = country_stats["Life satisfaction"].values m, b = np.polyfit(X, y, 1)
为什么Seaborn的regplot能正常工作
Seaborn的regplot会自动处理输入的二维数组,内部会将其展平为一维后再执行线性回归计算,因此不需要手动调整数组维度,结果自然贴合数据趋势。
修正后的完整代码
%matplotlib inline import matplotlib.pyplot as plt import pandas as pd import numpy as np # 原代码遗漏numpy导入 country_stats = pd.read_csv("../data/country_stats.csv") # 使用一维数组存储特征和标签 X = country_stats["GDP per capita"].values y = country_stats["Life satisfaction"].values country_stats.plot(kind='scatter', x="GDP per capita", y='Life satisfaction') plt.axis([0,60000,0,10]) # 拟合线性回归线 m, b = np.polyfit(X, y, 1) plt.plot(X, m*X+b, color='red') plt.show()
内容的提问来源于stack exchange,提问作者John Conor
相关产品推荐
相关产品推荐

