You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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是一维数组,你可以通过两种方式修正:

  1. 把二维数组展平:
m, b = np.polyfit(X.flatten(), y.flatten(), 1)
  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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.21 20:36:23