首页
/ ML-For-Beginners:用 Scikit-learn 以四种方式构建回归模型——南瓜价格预测实战

ML-For-Beginners:用 Scikit-learn 以四种方式构建回归模型——南瓜价格预测实战

2026-09-05 19:23:52作者:盛欣凯Ernestine

本篇技术指南基于 ML-For-Beginners 课程第 2 周(回归)第 3 课 2-Regression/3-Linear/README.md,围绕美国南瓜价格数据集,完整讲解如何用 Scikit-learn 依次构建简单线性回归、多项式回归、含类目特征(one-hot 编码)的线性回归,以及最终的“多项式 + 全特征”组合模型。读完本篇,你将掌握最小二乘回归的数学原理、相关性分析、train_test_splitmake_pipeline 的标准训练流程,以及从 RMSE 到决定系数(R²)的模型评估方法,并能在 solution 笔记本 中复现每步的真实输出。

线性回归与多项式回归对比信息图

什么是线性回归:预测数值型目标的起点

线性回归(Linear Regression)用于预测数值型结果,例如房价、气温或销售额。它的核心思想是找到一条能最好地刻画“输入特征与输出关系”的直线。在 ML-For-Beginners 的回归课程中,本节聚焦于先理解概念,再扩展到更高级的回归技术——即线性回归与多项式回归,并用它们预测南瓜价格(每蒲式尔价格)如何随月份、品种等输入变化。

本节使用的数据集是南瓜价格数据集,数据文件位于 2-Regression/data/US-pumpkins.csv,来源于美国农业部 Specialty Crops Terminal Markets 的公开报告(数据公有领域,详见 2-Regression/README.md 的 Credits 说明)。数据经过预加载与清洗后,以“每蒲式尔价格”的形式存入一个新的 DataFrame,随本课程的 notebook.ipynb 文件一起提供。

数据准备:从原始 CSV 到 new_pumpkins

2-Regression/3-Linear/notebook.ipynb 中,数据清洗的具体步骤是:

import pandas as pd
import matplotlib.pyplot as plt
import numpy as np
from datetime import datetime

pumpkins = pd.read_csv('../data/US-pumpkins.csv')
  1. 只保留按蒲式尔计价的南瓜pumpkins = pumpkins[pumpkins['Package'].str.contains('bushel', case=True, regex=True)]
  2. 取需要的列PackageVarietyCity NameLow PriceHigh PriceDate
  3. 计算平均价price = (pumpkins['Low Price'] + pumpkins['High Price']) / 2
  4. 时间特征month = pd.DatetimeIndex(pumpkins['Date']).month,以及课程 README 中给出的“一年中的第几天”计算式:
day_of_year = pd.to_datetime(pumpkins['Date']).apply(lambda dt: (dt-datetime(dt.year,1,1)).days)
  1. 价格归一化到“每蒲式尔”:对 1 1/9 bushel 包装的报价除以 1.1(因为 1 1/9 蒲式尔 ≈ 1.1 蒲式尔),对 1/2 bushel 包装的报价乘以 2,最终得到 new_pumpkins 数据框。

new_pumpkins 的结构如下(节选自课程 README):

ID Month DayOfYear Variety City Package Low Price High Price Price
70 9 267 PIE TYPE BALTIMORE 1 1/9 bushel cartons 15.0 15.0 13.636364
71 9 267 PIE TYPE BALTIMORE 1 1/9 bushel cartons 18.0 18.0 16.363636
72 10 274 PIE TYPE BALTIMORE 1 1/9 bushel cartons 18.0 18.0 16.363636
73 10 274 PIE TYPE BALTIMORE 1 1/9 bushel cartons 17.0 17.0 15.454545
74 10 281 PIE TYPE BALTIMORE 1 1/9 bushel cartons 15.0 15.0 13.636364

加载数据的目的是向数据“提问”:什么时候买南瓜最划算?一箱微型南瓜能预期什么价格?应该买半蒲式尔篮装还是 1 1/9 蒲式尔盒装?

最小二乘回归线:为什么要“平方”误差

线性回归的目标是画出一条线,用来:

  • 展示变量间的关系
  • 做预测:判断一个新的数据点相对于这条线会落在什么位置。

最小二乘回归(Least-Squares Regression) 是画这类线的典型方法。所谓“最小二乘”,是指让模型总误差最小化的过程:对每个数据点,测量该点到回归线的垂直距离(称为残差 residual),然后对这些距离平方并求和,目标是找到使这个总和最小的那条线。

为什么要对残差平方?课程给出两个主要原因:

  1. 看幅度而非方向(Magnitude over Direction):我们希望 -5 的误差和 +5 的误差被同等对待,平方后所有值都变为正数;
  2. 惩罚离群点(Penalizing Outliers):平方对较大的误差给予更大权重,迫使回归线更靠近距离线较远的点。

🧮 数学说明:Y = a + bX

最佳拟合线(line of best fit)可以用简单线性回归方程表示:

Y = a + bX
  • X 是“解释变量”(explanatory variable),Y 是“因变量”(dependent variable);
  • 直线的斜率是 ba 是 y 轴截距,即 X = 0Y 的取值。

结合南瓜数据的问题“按月份预测每蒲式尔南瓜价格”来理解这个方程:先计算斜率 b,斜率的计算又依赖于截距(即 X = 0Y 的所在位置)。截距与斜率确定后,代入新的 X 即可计算出对应的 Y——课程原文用它来演示“如果你看到的均价在 4 美元左右,那大概率是四月的南瓜”。

相关性:训练回归模型前的必经检查

在训练模型之前,需要理解 相关系数(Correlation Coefficient):用散点图可以快速直观地判断它——数据点大致沿一条直线分布说明相关性高,数据点在 X 与 Y 之间四处散落则相关性低。一个良好的线性回归模型,应该用最小二乘法拟合出相关系数较高(更接近 1 而非 0)的回归线。

先用 Pandas 的 corr 函数检查整体相关性:

print(new_pumpkins['Month'].corr(new_pumpkins['Price']))
print(new_pumpkins['DayOfYear'].corr(new_pumpkins['Price']))

Month 的相关性约为 -0.15,按 DayOfYear 的约为 -0.17——相关性很小,说明单用时间特征很难线性解释价格。但这可能掩盖了另一个更重要的关系:价格数据似乎按不同南瓜品种聚成了不同的簇。

按品种分色散点:发现真正的价格驱动因素

通过给 scatter 绑图函数传入 ax 参数,可以把所有品种画在同一张图上:

ax=None
colors = ['red','blue','green','yellow']
for i,var in enumerate(new_pumpkins['Variety'].unique()):
    df = new_pumpkins[new_pumpkins['Variety']==var]
    ax = df.plot.scatter('DayOfYear','Price',ax=ax,c=colors[i],label=var)

不同品种南瓜的 DayOfYear 与 Price 散点图(按品种着色)

分组均值柱状图进一步印证了结论:

new_pumpkins.groupby('Variety')['Price'].mean().plot(kind='bar')

南瓜品种平均价格柱状图

从图中可以判断:品种对整体价格的影响大于销售日期本身

聚焦单一品种:PIE TYPE 的时间相关性

先把范围缩小到一个品种——PIE TYPE(派用南瓜):

pie_pumpkins = new_pumpkins[new_pumpkins['Variety']=='PIE TYPE']
pie_pumpkins.plot.scatter('DayOfYear','Price')

此时计算 PriceDayOfYear 的相关性,得到约 -0.27(仓库 solution 笔记本 中该单元格的实际输出为 -0.2669192282197318),意味着训练一个预测模型是有意义的。

训练线性回归模型前,必须确认数据是干净的。线性回归对缺失值处理不好,因此合理的做法是去掉所有空单元格:

pie_pumpkins.dropna(inplace=True)
pie_pumpkins.info()

另一种处理方式是用对应列的均值填充这些空值。

简单线性回归:用 Scikit-learn 训练第一个模型

训练线性回归模型使用 Scikit-learn 库:

from sklearn.linear_model import LinearRegression
from sklearn.metrics import mean_squared_error
from sklearn.model_selection import train_test_split

首先把输入值(特征)和期望输出(标签)拆成独立的 numpy 数组:

X = pie_pumpkins['DayOfYear'].to_numpy().reshape(-1,1)
y = pie_pumpkins['Price']

注意这里对输入做了 reshape(-1,1):Scikit-learn 的 Linear Regression 期望二维数组作为输入,每一行对应一个输入特征向量。由于这里只有一个输入特征,需要把数组整形成 N×1(N 为数据集大小)。

然后划分训练集与测试集,以便在训练后验证模型:

X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=0)

其中 test_size=0.2 表示 20% 数据用作测试集,random_state=0 固定随机种子保证结果可复现。

训练真正的线性回归模型只需要两行代码——定义 LinearRegression 对象并调用 fit

lin_reg = LinearRegression()
lin_reg.fit(X_train,y_train)

fit 之后的 LinearRegression 对象包含回归的全部系数,可通过 .coef_ 属性访问。本例只有一个系数,约为 -0.017,意味着价格随时间略有下降,大约每天 2 美分。截距可通过 lin_reg.intercept_ 访问,约为 21,表示年初的价格水平。

评估:RMSE 与决定系数

用测试集上的预测值衡量模型精度,采用均方根误差(RMSE)——所有“期望值与预测值之差的平方”的均值的平方根:

pred = lin_reg.predict(X_test)

rmse = np.sqrt(mean_squared_error(y_test,pred))
print(f'RMSE: {rmse:3.3} ({rmse/np.mean(pred)*100:3.3}%)')

误差约为 2 美元(约 17%),不算理想。另一个模型质量指标是决定系数(coefficient of determination,即 R²)

score = lin_reg.score(X_train,y_train)
print('Model determination: ', score)

决定系数为 0 表示模型完全没利用输入数据,退化成了“最差线性预测器”(即直接输出结果均值);为 1 表示可以完美预测所有输出。本例约为 0.06,相当低——说明单靠时间这一特征,线性模型解释力很弱。

把测试数据与回归线一起画出来,可以更直观地看到效果:

plt.scatter(X_test,y_test)
plt.plot(X_test,pred)

多项式回归:拟合非线性价格波动

除了直线关系,有些变量间的关系无法用平面或直线表示——南瓜价格可能会随季节波动,直线未必是合适的刻画。多项式回归(Polynomial Regression)通过生成曲线来更好地拟合非线性数据。

在本例中,把 DayOfYear 的平方项加入输入数据,就能用一条抛物线去拟合数据,这条抛物线会在一年中的某个点取得最小值。

Scikit-learn 提供了 pipeline API 把多个数据处理步骤串起来。一个 pipeline 是一串 estimator(估计器) 的链。这里创建一个先添加多项式特征、再做回归训练的 pipeline:

from sklearn.preprocessing import PolynomialFeatures
from sklearn.pipeline import make_pipeline

pipeline = make_pipeline(PolynomialFeatures(2), LinearRegression())

pipeline.fit(X_train,y_train)

PolynomialFeatures(2) 表示包含输入数据的全部二阶多项式。单变量场景下就是 DayOfYear²;若有两个输入 X 和 Y,则会新增 XY。需要时也可以使用更高阶的多项式。

pipeline 的用法与原始 LinearRegression 对象完全一致:fit 之后用 predict 得到预测结果:

pred = pipeline.predict(X_test)

rmse = np.sqrt(mean_squared_error(y_test,pred))
print(f'RMSE: {rmse:3.3} ({rmse/np.mean(pred)*100:3.3}%)')

score = pipeline.score(X_train,y_train)
print('Model determination: ', score)

绘制平滑的近似曲线时,用 np.linspace 生成一段均匀取值范围,而不是直接在不有序的测试数据上连线(那样会画出锯齿线):

X_range = np.linspace(X_test.min(), X_test.max(), 100).reshape(-1,1)
y_range = pipeline.predict(X_range)

plt.scatter(X_test, y_test)
plt.plot(X_range, y_range)

多项式回归拟合曲线与测试散点

多项式回归得到的 RMSE 略低、决定系数略高,但改善并不显著(solution 笔记本 中该模型的实际输出为 RMSE 2.73(17.0%)、determination 约 0.076)——必须引入其他特征。顺带一提,从图上可以看到最低南瓜价格出现在万圣节前后,这与节日市场需求的变化是吻合的。

类目特征:one-hot 编码让品种进入模型

理想情况下,我们希望用同一个模型预测不同南瓜品种的价格。但 Variety 列与 Month 这类列不同,它包含非数值值——这样的列称为类目(categorical)特征,需要先转换为数值形式,即编码。有两种常见做法:

  • 简单数值编码:建立一张品种表,用品种名在表中的索引替换品种名。这对线性回归并不是好主意——线性回归会把索引的数值当作真实数字,乘以某个系数加入结果,而索引与价格之间显然不存在线性关系(即使刻意规定索引的排列顺序也不行)。
  • One-hot 编码:把 Variety 列替换为 4 列(每种品种一列),对应行属于该品种时取 1,否则取 0。这样线性回归中会有四个系数,每个品种一个,分别负责该品种的“起步价格”(更准确说是“额外价格”)。

one-hot 编码一行代码即可完成:

pd.get_dummies(new_pumpkins['Variety'])

输出类似:

ID FAIRYTALE MINIATURE MIXED HEIRLOOM VARIETIES PIE TYPE
70 0 0 0 1
71 0 0 0 1
... ... ... ... ...
1738 0 1 0 0
1739 0 1 0 0
1740 0 1 0 0
1741 0 1 0 0
1742 0 1 0 0

只用品种训练线性回归

用 one-hot 编码后的品种作为输入,只需正确初始化 Xy

X = pd.get_dummies(new_pumpkins['Variety'])
y = new_pumpkins['Price']

其余代码与前面的线性回归训练流程完全相同。实测结果:均方误差(RMSE)与之前相近,但决定系数大幅提升到约 77%solution 笔记本 中的对应输出为 Mean error: 5.24 (19.7%)Model determination: 0.774085281105197)——这验证了前面散点图的分析:品种才是价格的主要驱动因素。

组合数值特征与类目特征

要得到更精确的预测,可以把更多类目特征和数值特征(如 MonthDayOfYear)一起纳入。用 join 拼成一个大特征数组:

X = pd.get_dummies(new_pumpkins['Variety']) \
        .join(new_pumpkins['Month']) \
        .join(pd.get_dummies(new_pumpkins['City'])) \
        .join(pd.get_dummies(new_pumpkins['Package']))
y = new_pumpkins['Price']

这里把 City(城市)和 Package(包装规格)也加入了特征,得到 RMSE 2.84(10.5%)、决定系数 0.94(与 solution 笔记本输出 Mean error: 2.84 (10.5%)Model determination: 0.9401096672643048 一致)。

把所有技巧合起来:全特征多项式模型

要得到最好的模型,把上面“one-hot 编码类目特征 + 数值特征”的数据与多项式回归组合起来。课程 README 给出的完整代码如下:

# set up training data
X = pd.get_dummies(new_pumpkins['Variety']) \
        .join(new_pumpkins['Month']) \
        .join(pd.get_dummies(new_pumpkins['City'])) \
        .join(pd.get_dummies(new_pumpkins['Package']))
y = new_pumpkins['Price']

# make train-test split
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=0)

# setup and train the pipeline
pipeline = make_pipeline(PolynomialFeatures(2), LinearRegression())
pipeline.fit(X_train,y_train)

# predict results for test data
pred = pipeline.predict(X_test)

# calculate RMSE and determination
rmse = mean_squared_error(y_test, pred, squared=False)
print(f'RMSE: {rmse:3.3} ({rmse/pred.mean()*100:3.3}%)')

score = pipeline.score(X_train,y_train)
print('Model determination: ', score)

版本提示:mean_squared_errorsquared=False 参数在较新版本的 scikit-learn 中直接返回 RMSE;在旧版本中该函数返回 MSE,需要像前文那样自行套一层 np.sqrt(...)(solution 笔记本中正是这么写的)。两种写法按本机 scikit-learn 版本择一即可。

最终效果:决定系数接近 97%,RMSE = 2.23(约 8% 的预测误差),与 solution 笔记本 的最终单元格输出 Mean error: 2.23 (8.25%)Model determination: 0.9652870784724543 完全一致。

四个模型对比总表

课程 README 汇总了本课时依次构建的四个回归模型的实测指标:

模型 RMSE 决定系数
DayOfYear 线性回归 2.77(17.2%) 0.07
DayOfYear 多项式回归 2.73(17.0%) 0.08
Variety 线性回归 5.24(19.7%) 0.77
全特征线性回归 2.84(10.5%) 0.94
全特征多项式回归 2.23(8.25%) 0.97

这张表清晰展示了本课程的完整技术路线:单变量线性(解释力弱)→ 多项式(略改善)→ one-hot 类目特征(R² 跃升到 0.77)→ 全特征线性(R² 0.94)→ 全特征多项式(R² 0.97)。特征工程(尤其是正确的类目编码)对本例模型质量的贡献远超单纯提升模型复杂度。

相关性的再验证

✅ 动手练习(Challenge):在本课程笔记本中尝试若干不同的变量,观察相关性与模型精度之间的对应关系——比如 DayOfYearPrice 在全体数据上相关性只有 -0.17,而在 PIE TYPE 子集上达到 -0.27,模型的决定系数也随之从“几乎无用”提升到“有一定解释力”。

学习资源与后续方向

  • 完整可运行实现:数据清洗与建模全流程在 2-Regression/3-Linear/solution/notebook.ipynb 中给出了带真实输出的完整版本;学生练习版为 2-Regression/3-Linear/notebook.ipynb(仅含数据加载与散点图部分)。
  • R 语言版本:本课同时提供 R 实现,可查阅 2-Regression/3-Linear/solution/R/lesson_3.html
  • 作业(Assignment):按 assignment.md 的要求,自选一个数据集(或使用 Scikit-learn 内置数据集)构建一个全新的线性/多项式回归模型,并在笔记本中说明技术选型理由、展示模型精度;若精度不理想,还需解释原因。评分标准以“完整、文档齐全的笔记本”为优秀档。
  • 进阶阅读:本课时聚焦线性回归,其他重要的回归类型包括 Stepwise、Ridge、Lasso 与 ElasticNet,值得进一步学习;课程推荐以斯坦福的统计学习(Statistical Learning)课程为延伸材料。
  • 课程地图:本节是“回归”四课中的第 3 课(1. 工具、2. 数据管理、3. 线性与多项式回归、4. 逻辑回归),下一章将学习用逻辑回归确定类别,详见 2-Regression/README.md
登录后查看全文
热门项目推荐
相关项目推荐