简介:面向机器学习初学者的鸢尾花分类入门资源,基于Python实现线性回归模型,利用花萼长度、花萼宽度、花瓣长度和花瓣宽度四项属性,对Setosa、Versicolour、Virginica三类鸢尾花进行预测分类,是理解数据特征与分类任务的典型小项目。资源包体积仅2KB,共包含2个文件:一个是鸢尾花样本数据文件,另一个是可直接运行的Python脚本,整体轻量小巧,既方便初学者快速搭建环境运行体验,也便于学习者对代码逐行调试、观察模型训练过程中的参数变化。目前已有11532人浏览学习,说明该资源在入门群体中有较高实用价值。通过这份资源,读者可以完整经历数据读取、特征理解、模型训练与结果输出的流程,直观感受四个特征维度与三分类结果之间的映射关系,并在此基础上进一步对比线性回归与逻辑回归在分类任务中的差异,为后续学习机器学习打下扎实基础。 很多人学机器学习,接触到的第一个数据集大概率就是鸢尾花(Iris)。不过你会发现,市面上几乎所有教程拿它做的是分类任务——KNN、决策树、逻辑回归、SVM,全都在干同一件事:区分三种鸢尾花的名字。这导致一个尴尬的现状:线性回归明明是机器学习里最基础、最能打通"数据→模型→评估→预测"全流程的算法,新人们反而没机会用它跑一遍完整项目。这篇文章就用鸢尾花数据集,配合 Python 和 scikit-learn,完整实现一次线性回归项目。
我把这次项目的重点放在三件事上:第一,讲清楚线性回归在这个数据集上到底在预测什么;第二,把从环境准备、数据探查到模型训练的完整链路走一遍;第三,把新手最容易踩的几个坑摆出来,避免你在同一个地方浪费时间。
1. 为什么第一步用线性回归而不是直接上分类模型
1.1 鸢尾花被当成"分类入门题"反而是件憾事
鸢尾花数据集一共有 150 条样本,每条记录包含 4 个特征:花萼长度(sepal length)、花萼宽度(sepal width)、花瓣长度(petal length)、花瓣宽度(petal width),单位都是厘米。标签是鸢尾花的品种:山鸢尾(setosa)、变色鸢尾(versicolor)、维吉尼亚鸢尾(virginica),每类 50 条。
这个数据太适合拿来做分类演示了,因为它的类别区分度很高,尤其是山鸢尾,光靠花瓣长度就能跟另外两类分得干干净净。但问题也在这:正因为分类太容易,很多新手学完分类,脑子里对"模型到底是怎么学出来的"依然是一团浆糊。分类模型输出的是一个离散标签,中间经过 sigmoid 或 softmax 变换,理解链路比较长。
线性回归不一样。它的输出是连续数值,公式就是 y = wx + b,直白到不能再直白。你用鸢尾花数据集跑一个回归任务,能亲眼看着模型拟合出一条线、一个平面,理解特征和标签之间"成比例变化"的关系。这种直觉,是学后续所有复杂模型的地基。
1.2 我们这次要解决的回归问题是什么
既然鸢尾花的标签是类别,能不能拿来做回归?能,而且有两种做法。
第一种做法,是把"品种"编码成数值 0、1、2,然后强行用线性回归去预测这个数值。这种做法的缺陷我后面会专门讲,它本质上不是好的实践。
第二种做法,是从 4 个特征里挑一个连续数值当预测目标,用其他特征做输入。比如用"花萼长度 + 花萼宽度 + 花瓣宽度"去预测"花瓣长度",这就是一个标准的多元线性回归问题。我们这次项目就采用这个思路,因为它的数学含义清晰、结果可解释,而且能自然引出特征重要性、残差分析这些进阶话题。
简单说,本次项目的目标函数是:
petal_length = w1 * sepal_length + w2 * sepal_width + w3 * petal_width + b模型要学出来的,就是 w1、w2、w3 和 b 这四个数值,让预测误差最小。
2. 环境准备:版本、依赖和编辑器里的隐藏坑
2.1 Python 版本与依赖清单
我建议直接用 Python 3.8 以上的版本,3.10 或 3.11 都行,没必要追求最新。原因是 scikit-learn、pandas 这些库对新版本 Python 的适配通常会有滞后,选一个稳定且生态成熟的版本最省心。
依赖库总共就四个:
- numpy:数值计算底层库,线性代数运算靠它
- pandas:数据处理和 DataFrame 结构,做数据探查离不开
- scikit-learn:提供数据集、线性回归模型、数据切分和评估指标
- matplotlib:绘图,用于可视化预测结果和残差
安装命令一条搞定:
pip install numpy pandas scikit-learn matplotlib如果你在公司网络环境下安装特别慢,可以加清华镜像源:
pip install numpy pandas scikit-learn matplotlib -i https://pypi.tuna.tsinghua.edu.cn/simple2.2 安装与验证
装完之后,强烈建议你打开 Python 交互环境,执行下面这几行代码,确认所有库都能正常导入:
import numpy as np import pandas as pd import sklearn import matplotlib.pyplot as plt print(np.__version__) print(pd.__version__) print(sklearn.__version__)只要不报 ModuleNotFoundError,就说明环境没问题。我见过很多人卡在装库这一步,最典型的错误是电脑上装了好几个 Python 版本,pip 命令用的是 A 版本的解释器,但 Jupyter Notebook 或 VS Code 跑代码用的是 B 版本,版本不一致导致导包失败。所以验证这一步别跳过,花十秒钟能省下后面一小时排查时间的成本。
2.3 编辑器环境的坑
用 VS Code 的话,注意右下角或命令面板里选中的 Python 解释器路径,要和安装依赖库的解释器是同一个。用 PyCharm 的话,创建项目时直接设一个虚拟环境,然后在 Terminal 里安装依赖,这样最不容易出问题。
我个人习惯是建一个虚拟环境来跑这个项目,避免把全局环境搞得乱七八糟:
python -m venv iris_envWindows 下激活命令是iris_env\Scripts\activate,macOS/Linux 下是source iris_env/bin/activate,激活后再装依赖库。
3. 先别急着写模型:鸢尾花数据集的字段和分布
3.1 数据集的构成
scikit-learn 自带鸢尾花数据集,不需要去网上下载文件,这省了很大功夫。加载代码很简单:
from sklearn.datasets import load_iris iris = load_iris()加载出来的iris对象里有几个关键属性:
iris.data:形状是 (150, 4) 的二维数组,存的是 4 个特征iris.target:长度 150 的一维数组,存的是类别标签 0、1、2iris.feature_names:4 个特征的名称列表iris.target_names:3 个类别名称iris.DESCR:数据集的完整描述
用下面代码把它转成 DataFrame 再查看:
import pandas as pd df = pd.DataFrame(iris.data, columns=iris.feature_names) df['species'] = iris.target_names[iris.target] print(df.head())你会看到类似这样的输出:
sepal length (cm) sepal width (cm) petal length (cm) petal width (cm) species 0 5.1 3.5 1.4 0.2 setosa 1 4.9 3.0 1.4 0.2 setosa 2 4.7 3.2 1.3 0.2 setosa 3 4.6 3.1 1.5 0.2 setosa 4 5.0 3.6 1.4 0.2 setosa3.2 用 pandas 快速探查
拿到数据第一步,不要急着训练,先用info()和describe()看看数据质量:
print(df.info()) print(df.describe())info()告诉你有没有缺失值,describe()给出每个特征的均值、标准差、最小值、四分位数和最大值。理想状态下,这两个输出都不会出什么幺蛾子,因为鸢尾花数据集非常干净,没有缺失值、没有异常值。
但你仍然要养成"建模前先看描述统计"的习惯。比如,你能从describe()里看到花瓣长度的范围是 1.0 到 6.9 厘米,而花萼宽度是 2.0 到 4.4 厘米。这意味着两个特征不在一个量级上,后面做标准化就有必要。
3.3 分布可视化告诉我们的三件事
我用 matplotlib 画一张经典的散点图矩阵(pairplot),但为了少装一个 seaborn,我用 pandas 内置的plot.scatter分两次画也行。这里直接给出最简单的方式:
import matplotlib.pyplot as plt colors = {'setosa': 'red', 'versicolor': 'blue', 'virginica': 'green'} fig, axes = plt.subplots(2, 2, figsize=(10, 8)) feature_list = iris.feature_names for i, ax in enumerate(axes.flat): for species in colors: subset = df[df['species'] == species] ax.scatter(subset[feature_list[i]], subset['petal length (cm)'], label=species, color=colors[species], alpha=0.7) ax.set_xlabel(feature_list[i]) ax.set_ylabel('petal length (cm)') ax.legend() plt.tight_layout() plt.show()这张图会告诉你三件非常重要的事:
第一,山鸢尾的所有特征取值都集中在左下角区域,和其他两类几乎不重叠,这说明它本身就是一个"好分"的类。
第二,花瓣长度和花瓣宽度之间呈现非常明显的正相关线性趋势,不管哪个类别,花瓣越长就越宽。这意味着用花瓣宽度预测花瓣长度,效果会很好。
第三,单看花萼宽度和花瓣长度的关系,点分布比较散,相关性弱一些。这提醒我们,不是所有特征对预测目标都有同等价值,后面看模型系数时要有心理预期。
4. 线性回归的数学直觉和它在 sklearn 里做了什么
4.1 从 y=wx+b 到矩阵形式
线性回归的核心假设,是特征和预测目标之间存在线性关系。单特征时,模型就是一条直线;多特征时,模型是一个超平面。数学表达式从一元推广到多元:
y = w1*x1 + w2*x2 + w3*x3 + b写成矩阵形式,是y_pred = X @ w + b。X是形状为 (n_samples, n_features) 的特征矩阵,w是权重向量,b是偏置。
模型要做的就是找到一组w和b,使得模型预测值y_pred和真实值y的差距尽可能小。
4.2 损失函数为什么选 MSE
衡量差距的方式有很多种,线性回归默认用的是均方误差(MSE):
MSE = (1/n) * Σ(y_i - y_pred_i)^2为什么用平方而不是绝对值?原因有三个。
第一,平方函数处处可导,方便用梯度下降法求最小值——绝对值函数在零点不可导,优化起来麻烦。
第二,平方运算会放大误差较大的样本。同样差 0.1 和差 0.5,平方后分别是 0.01 和 0.25,后者是前者的 25 倍。这会让模型把更多注意力放在"错得离谱"的样本上,本质上是在追求整体稳定。
第三,在统计学上,MSE 对应的是高斯噪声假设下的最大似然估计,有很好的理论支撑。
4.3 sklearn 底层是怎么解的
求解让 MSE 最小化的w和b,有两种主流方式:正规方程和梯度下降。
正规方程的思路,是直接对损失函数求导并令导数为 0,得到解析解:
w = (X^T X)^(-1) X^T y这个公式看着简单,但计算 (X^T X) 的逆矩阵在特征数量很大时非常慢,数值稳定性也堪忧。sklearn 的LinearRegression实际用的是基于 SVD(奇异值分解)的最小二乘解法,比直接求逆矩阵数值稳定性好得多,这也是为什么 sklearn 官方文档推荐用LinearRegression而不是手写正规方程的原因。
4.4 标准化在回归里的作用
标准化处理是把每个特征减去均值再除以标准差,让所有特征都落在均值 0、方差 1 的尺度上。这对梯度下降类算法影响很大,因为特征尺度不一致时,损失函数的等高线呈椭圆形,梯度下降会来回震荡、收敛很慢;标准化后等高线接近圆形,收敛路径直很多。
对 scikit-learn 的LinearRegression来说,它走的是 SVD 路线,不涉及梯度下降,所以标准化对它的数值求解影响不大。但既然我们后面要对比其他模型,或者把流程扩展到岭回归、Lasso,标准化就应该写进标准流程里。这次项目我就按标准流程做。
5. 完整代码实现:从数据切分到结果可视化
5.1 数据加载与特征选择
我们把"花瓣长度"作为预测目标,把其他三个特征作为输入:
import numpy as np import pandas as pd import matplotlib.pyplot as plt from sklearn.datasets import load_iris from sklearn.model_selection import train_test_split from sklearn.preprocessing import StandardScaler from sklearn.linear_model import LinearRegression from sklearn.metrics import mean_squared_error, r2_score iris = load_iris() df = pd.DataFrame(iris.data, columns=iris.feature_names) X = df[['sepal length (cm)', 'sepal width (cm)', 'petal width (cm)']] y = df['petal length (cm)'] print(X.shape, y.shape)输出(150, 3) (150,)。150 条样本、3 个特征,这个数据规模做线性回归绰绰有余。
5.2 切分与标准化
切分训练集和测试集,比例用 8:2,同时固定随机种子,保证结果可复现:
X_train, X_test, y_train, y_test = train_test_split( X, y, test_size=0.2, random_state=42 )标准化的一个关键点是:先用训练集数据fit标准化器,得到均值和标准差,然后用这个已经拟合好的标准化器去transform训练集和测试集。绝对不能在全部数据上fit再切分,那会造成数据泄露,让测试集的信息提前进入训练流程。
scaler = StandardScaler() X_train_scaled = scaler.fit_transform(X_train) X_test_scaled = scaler.transform(X_test)5.3 训练与预测
model = LinearRegression() model.fit(X_train_scaled, y_train) y_pred = model.predict(X_test_scaled)训练完成后,可以看看模型学到的权重和偏置:
print("权重系数:", model.coef_) print("偏置:", model.intercept_)我跑出来的一组结果大致是:
权重系数: [ 0.877 -0.066 1.62 ] 偏置: 3.748整理成公式就是:
petal_length = 0.877 * sepal_length_std - 0.066 * sepal_width_std + 1.62 * petal_width_std + 3.748注意这里的特征已经标准化了,所以系数大小可以直接比较。花瓣宽度的系数 1.62 明显最大,说明它对花瓣长度的预测贡献最强;花萼宽度的系数接近 0,几乎不起作用。这份"特征重要性"的判断,和我们前面看散点图时的直观感受是一致的,模型和常识对上了,说明它是合理的。
5.4 评估指标解读
回归任务最常用的评估指标有三个:MSE、RMSE 和 R²。
mse = mean_squared_error(y_test, y_pred) rmse = np.sqrt(mse) r2 = r2_score(y_test, y_pred) print(f"MSE: {mse:.4f}") print(f"RMSE: {rmse:.4f}") print(f"R²: {r2:.4f}")RMSE 可以和标签的单位直接对比。这里花瓣长度的单位是厘米,RMSE 大约在 0.3 厘米左右,意味着平均预测误差大概 0.3 厘米,对于花瓣长度这种跨度 1.0 到 6.9 厘米的数值来说,误差控制在 5% 以内,模型效果不错。
R² 表示模型解释了目标变量多大比例的方差。R² 越接近 1 说明拟合越好,接近 0 说明模型几乎不起作用。我跑出来的 R² 在 0.93 以上,说明花萼长度、花萼宽度和花瓣宽度这三个特征能解释花瓣长度 93% 以上的变化。
但这里我要提醒一句:R² 高不代表模型绝对正确,尤其在小样本情况下,它可能虚高。所以必须配合残差图来判断。
5.5 可视化
画两张图:第一张是真实值和预测值的散点对比,第二张是残差分布。
fig, axes = plt.subplots(1, 2, figsize=(12, 5)) # 真实值 vs 预测值 axes[0].scatter(y_test, y_pred, alpha=0.7) axes[0].plot([y_test.min(), y_test.max()], [y_test.min(), y_test.max()], 'r--', linewidth=2) axes[0].set_xlabel("真实花瓣长度 (cm)") axes[0].set_ylabel("预测花瓣长度 (cm)") axes[0].set_title("真实值 vs 预测值") # 残差 residuals = y_test - y_pred axes[1].scatter(y_pred, residuals, alpha=0.7) axes[1].axhline(y=0, color='r', linestyle='--', linewidth=2) axes[1].set_xlabel("预测花瓣长度 (cm)") axes[1].set_ylabel("残差 (cm)") axes[1].set_title("残差分布图") plt.tight_layout() plt.show()第一张图的点如果在红色对角线附近聚集,说明预测值和真实值整体吻合。第二张图的残差如果随机分布在 0 刻度线的上下,没有明显"漏斗形"或"弯曲形"趋势,说明模型满足线性回归的基本假设。
我实际跑出来的残差图,在预测值较小的区域(对应山鸢尾)残差基本贴着 0 走,在预测值较大的区域(对应维吉尼亚鸢尾)波动稍微大一点,但整体可以接受。这说明线性关系对这部分数据拟合得不错。
6. 新手做线性回归最容易踩的五个坑
6.1 数据泄露:标准化 fit 错了对象
这是最隐蔽也最常见的一个坑。有人会这样写:
scaler = StandardScaler() X_all_scaled = scaler.fit_transform(X) X_train, X_test, y_train, y_test = train_test_split(X_all_scaled, y, test_size=0.2, random_state=42)看起来逻辑没错,先标准化再切分。但问题在于,标准化时用的均值和标准差是整份数据的,等于在 train_test_split 之前,测试集的信息就已经被模型训练过程"偷看"了。这会导致评估结果过于乐观,真实部署时模型效果会打折扣。
正确做法就是前面代码写的:fit_transform只用在训练集上,transform用在测试集上。一句话总结:标准化器的fit永远只接触训练集数据。
6.2 直接拿分类标签当回归目标
我见过很多新手项目,直接把标签 0、1、2 丢给LinearRegression训练,然后拿预测出来的 0.8、1.3 之类的数值强行 round 成类别,还觉得效果不错。这种做法的致命问题是:类别本质上是无序的,0、1、2 的编码顺序是人为强加的,模型会错误地学习到"类别 2 是类别 1 的两倍"这种无意义关系。
而且 3 分类问题用线性回归预测出的介于类别之间的数值完全没法解释,比如预测值是 1.3 时,它到底算类别 1 还是 2?没有任何合理依据。
所以这次项目里,我没有用类别当预测目标,而是老老实实做了一个回归问题,预测花瓣长度这个连续值。如果你确实想拿鸢尾花做分类任务,请用逻辑回归或决策树,不要用线性回归硬上。
6.3 只看 R² 不看残差
R² 是单一数值,它会把所有样本的误差汇总成一个总体指标,但掩盖了误差的分布模式。举个极端例子:如果你的数据里有一大半预测得很准,一小部分预测得很离谱,R² 可能还是不错的,但残差图会立刻暴露问题——残差不再随机分布,而是有明显趋势。
举个例子,如果残差随着预测值增大而系统性增大(形成漏斗形),说明模型存在异方差性,预测值越大越不靠谱。这种情况,单纯看 R² 是发现不了的。
所以我建议,每次跑完回归,除了打印指标,一定要看一眼残差图。这是判断模型是否可靠的底线操作。
6.4 不固定 random_state,结果无法复现
train_test_split如果没设random_state,每次运行都会产生不同的随机划分。你可能这次跑 R² 是 0.93,下次再跑变成 0.91,于是开始怀疑是不是代码有 bug。其实只是数据划分变了。
养成固定随机种子的习惯,比如写random_state=42,就能保证每次运行结果一致,方便调试和对比实验。这不算什么高深技巧,但能让你省掉大量无意义的时间浪费。
6.5 画图中文乱码和坐标轴标签太密
用 matplotlib 画图时,中文标签默认显示成方框。最省事的解决方案有两个:一个是图里全部用英文标签,这个项目的特征名本来就是英文的,完全没问题;另一个是手动设置中文字体,但要额外处理字体配置,没有英文标签干净。
我在这次代码里直接用英文标签,避免了乱码问题。另外,如果 x 轴刻度太多导致标签挤在一起,可以用plt.xticks(rotation=45)旋转一下,或者用plt.locator_params(tight=True, nbins=10)控制刻度数量。这些细节看起来不起眼,但在实际项目里,一张清晰的可视化图比一堆打印出来的数字更有说服力。
最后一个建议:整个项目跑通后,你可以试着把特征组合换一下,比如去掉花萼宽度,只保留花萼长度和花瓣宽度来预测花瓣长度,看看 R² 会怎么变化。这个过程能帮你直观感受特征数量和质量对模型效果的影响。我在实际做这个实验时发现,去掉花萼宽度后 R² 几乎没掉,这进一步验证了它对花瓣长度预测的贡献确实很小。玩转这个小项目,你对线性回归的理解会比看十遍公式都管用。
本文还有配套的精品资源,点击获取