news 2026/9/9 15:47:16

机器学习入门实战:用Scikit-learn实现鸢尾花KNN分类

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
机器学习入门实战:用Scikit-learn实现鸢尾花KNN分类

先交代一下背景:这篇内容是我自己“机器学习进阶系列”的第三篇。前面两篇我们聊了机器学习到底在干嘛、常用术语是什么意思,到这一篇,终于要写第一行能跑的代码了。我特意选了鸢尾花分类这个经典到不能再经典的例子,不是因为花样多,而是因为它足够小、足够干净、足够让一个新手在半小时内走完“数据加载→模型训练→评估预测”的完整流程。就算你之前完全没碰过 Python 和 Scikit-learn,跟着这篇一步步敲,也能把整个流程跑通,而且能真正理解每一行在干什么,而不是复制粘贴完就关掉。

鸢尾花数据集是一个 150 条样本、4 个特征、3 个类别的标准分类数据集。它解决了什么问题?说白了,就是让我们用花的萼片长度、萼片宽度、花瓣长度、花瓣宽度这 4 个数值,去判断这朵花属于 Setosa、Versicolor 还是 Virginica。数据量不大,训练极快,非常适合做机器学习的“Hello World”。这篇我尽量按自己踩坑后的经验来写,不整虚的,把每一步的原理解释清楚,再把容易出错的地方标出来。

1. 机器学习项目的基本套路:从数据到模型的全流程

很多人学机器学习最大的困惑不是某个算法不会,而是拿到一个问题之后不知道该先干什么、后干什么。其实标准流程就那么几步:理解问题、准备数据、划分训练集和测试集、选择模型、训练模型、评估模型、调参优化。鸢尾花分类恰恰能把这几步全部走一遍,而且每一步的反馈都非常快,不会让你在某个环节卡太久。

我先用生活化的方式说一下这个流程。你想想你做饭:先看冰箱里有什么菜(数据探索),然后洗菜切菜(数据预处理),再把菜分两堆,一堆今晚吃(训练集),一堆明天吃(测试集),接着按菜谱炒(模型训练),最后尝一口看看咸淡(模型评估),咸了就加点盐再炒(调参)。机器学习本质上就是这套流程的程序化版本,只不过“尝咸淡”变成了计算准确率、看混淆矩阵。

1.1 为什么入门首选鸢尾花数据集

初学者最容易陷入的误区是:一上来就搞真实业务数据,比如用户行为日志、电商订单表,结果光是清洗数据就搞了一个月,还没看到模型长什么样就没耐心了。鸢尾花数据集好就好在它几乎不需要预处理:没有缺失值、没有异常值、特征全是数值型、类别是均衡的(每个类别恰好 50 条)。

另外一个原因是它只有 4 个特征,你可以在二维平面里直接画出来看分布。这一点极其重要,因为机器学习不是“把数据扔进模型就完事”,你需要直观理解数据长什么样、类别能否分得开。当你能用散点图看到 Setosa 和另外两类天然分离、Versicolor 和 Virginica 有部分重叠的时候,你就会明白为什么模型会出现某些“误判”,这种直观感受是拿大模型练手练不出来的。

1.2 分类问题的核心逻辑:让计算机学会“看特征下结论”

分类问题在机器学习里的定义很朴素:给定一组输入特征 X,预测一个离散的标签 y。放在鸢尾花语境下,X 就是 (150, 4) 的特征矩阵,y 就是 (150,) 的形状标签数组,每个元素取值是 0、1、2,分别对应该花的品种。

理解这个“X 和 y”的约定非常重要,因为你之后看所有 Scikit-learn 文档,都会看到 fit(X, y)、predict(X) 这样的签名。fit 的含义就是“用数据学习规律”,predict 的含义就是“用学到的规律去预测新数据”。这俩就像是练习和考试的关系:fit 是做练习题,predict 是上考场。你要是把练习题和考试卷混在一起,成绩自然虚高,这正是为什么后面要强调 train_test_split。

2. 环境准备与数据初探:跑起来之前先看懂数据

写第一行代码之前,先把环境准备好。别小看这一步,我见过太多人在环境上卡住,然后果断放弃。其实做机器学习入门,只需要装好 Python 3.8 以上的版本,然后用 pip 安装三个库就够了:scikit-learn、pandas、matplotlib。如果你用的是 Anaconda 发行版,那这三个库大概率已经预装好了,可以直接跳过安装步骤。

pip install scikit-learn pandas matplotlib

2.1 加载数据集并查看数据结构

Scikit-learn 内置了鸢尾花数据集,加载方式比从 CSV 文件读要简单得多,适合入门阶段集中精力理解建模逻辑。加载代码如下:

from sklearn.datasets import load_iris # 加载鸢尾花数据集 iris = load_iris() # 数据基本结构 print("特征矩阵形状:", iris.data.shape) # (150, 4) print("标签数组形状:", iris.target.shape) # (150,) print("类别名称:", iris.target_names) # ['setosa' 'versicolor' 'virginica'] print("特征名称:", iris.feature_names) # ['sepal length (cm)', 'sepal width (cm)', 'petal length (cm)', 'petal width (cm)']

这段代码的输出会告诉你三件事:样本量是 150,特征数是 4,类别数是 3。注意 iris.data 是一个 NumPy 二维数组,不是 DataFrame,所以你看不到列名,只有数值。如果想看得更舒服一点,可以用 pandas 包一层:

import pandas as pd df = pd.DataFrame(iris.data, columns=iris.feature_names) df['target'] = iris.target print(df.head()) print(df.describe())

head() 输出前 5 行,describe() 输出每个特征的均值、标准差、最小值、四分位数、最大值。这一步能帮你快速判断数据的量纲差异。你会发现花萼宽度(sepal width)的数值范围在 2.0 到 4.4 之间,花瓣长度(petal length)则在 1.0 到 6.9 之间,量纲差异不大。但在其他数据集中,如果某个特征的数值是几千,另一个是零点几,有些对距离敏感的算法(比如 KNN、SVM)就需要做标准化,这一点留到后面细说。

2.2 可视化:用散点图看看数据是否可分

一个很常见的初学者操作是:数据加载完立刻开始训练模型,完全不做可视化。但我想说,可视化这一步花不了两分钟,却能帮你建立对数据的直觉判断。核心问题是:这 3 类花能分开吗?用哪些特征分得最开?

下面这段代码选了花瓣长度和花瓣宽度两个特征,画一张散点图,并用不同颜色标记不同类别:

import matplotlib.pyplot as plt plt.figure(figsize=(8, 6)) scatter = plt.scatter( iris.data[:, 2], # 花瓣长度 iris.data[:, 3], # 花瓣宽度 c=iris.target, cmap='viridis', edgecolor='k', s=80 ) plt.colorbar(scatter, ticks=[0, 1, 2], label='品种') plt.xlabel('花瓣长度 (cm)') plt.ylabel('花瓣宽度 (cm)') plt.title('鸢尾花数据分布') plt.show()

跑完你就能看到,Setosa(图中颜色较深的一簇)在左下角,和另外两类完全分开;Versicolor 和 Virginica 虽然大致分成两片区域,但有少量边缘点互相重叠。这个发现直接影响了你能达到的上限:如果用花瓣长度加花瓣宽度两个特征组合,分类器已经把 Setosa 完美区分,剩下只需要尽量把 Versicolor 和 Virginica 的边界画对就行。

顺带说一句,如果你把花萼宽度和花萼长度画出来,会发现重叠区域非常大,这说明花萼特征区分度相对较弱。这种“先看数据再选特征”的习惯,比直接堆模型重要得多。

3. 第一行机器学习代码:数据划分与 KNN 模型训练

现在进入正题。所谓“第一行机器学习代码”,我把它定义为 model.fit(X_train, y_train) 这一行。为了让这一行能跑,我们要先完成数据划分。

3.1 为什么要划分训练集和测试集

先说结论:如果不划分,直接拿 150 条数据训练,再拿同一批数据评估,得到的准确率是虚高的,因为模型已经“见过”这些数据了。这相当于考试的时候把答案翻出来抄了一遍,考了 100 分,但换个题目就露馅。

正确的做法是把数据分成两部分:训练集用来让模型学习规律,测试集用来模拟“未来没见过的新数据”,只用于最终评估。在代码里,这行就能搞定:

from sklearn.model_selection import train_test_split X_train, X_test, y_train, y_test = train_test_split( iris.data, iris.target, test_size=0.2, random_state=42, stratify=iris.target )

每个参数都值得单独解释一下:

  • test_size=0.2 表示 20% 的数据划入测试集。由于样本量只有 150,测试集就是 30 条,训练集 120 条。
  • random_state=42 是随机种子。设成固定值后,每次运行代码划分结果完全一致,这样实验结果可复现,排除了“这次跑得好,下次跑得差”的随机干扰。42 没有什么特别含义,就是个习惯数值,你可以换 0、1、7 都行。
  • stratify=iris.target 是分层采样,按原始类别比例来划分。原本 3 类各 50 条,划分后每类在训练集里各占约 40 条、测试集里各占约 10 条。这个参数在类别不平衡的数据集上尤其重要,否则可能随机划分后测试集里缺了某个类别,评估结果就没有参考价值。

3.2 用 KNN 构建你的第一个模型

K近邻(K-Nearest Neighbors,简称 KNN)是我推荐新人尝试的第一个算法,核心思想非常朴素:“物以类聚,人以群分”。新样本进来,看看离它最近的 K 个训练样本属于哪些类别,多数票获胜。它不需要训练过程(严格说是惰性学习),所以代码特别清爽。

from sklearn.neighbors import KNeighborsClassifier # 创建模型,设置邻居数为 5 knn = KNeighborsClassifier(n_neighbors=5) # 训练模型 knn.fit(X_train, y_train) # 预测测试集 y_pred = knn.predict(X_test) # 输出预测结果 print("预测结果:", y_pred) print("真实标签:", y_test)

运行后,你会看到两个数组。预测结果是一组 0、1、2 的序列,真实标签是另一组 0、1、2 的序列。对一下,能对上多少个,就是后面要计算的准确率。

之前提到过 KNN 对距离敏感,所以这一步我还得再补充一句:因为鸢尾花 4 个特征的量纲都在厘米级别,差异不大,所以不标准化问题也不大。但如果换到其他数据,比如一个特征是“年龄”(20~60),另一个是“收入”(5000~50000),收入就会在距离计算中占据绝对主导地位,KNN 的效果就会崩。遇到量纲差距大的数据集,记得先做 StandardScaler 标准化。

3.3 评估模型:准确率、混淆矩阵与分类报告

预测完不评估等于白做。Scikit-learn 提供了多个评估函数,我建议初学者至少学会看三个:准确率(accuracy_score)、混淆矩阵(confusion_matrix)、分类报告(classification_report)。

from sklearn.metrics import accuracy_score, confusion_matrix, classification_report # 准确率:预测正确的样本数 / 总样本数 acc = accuracy_score(y_test, y_pred) print(f"测试集准确率: {acc:.2f}") # 混淆矩阵:每个类别的预测与真实对照 cm = confusion_matrix(y_test, y_pred) print("混淆矩阵:") print(cm) # 分类报告:精确率、召回率、F1 值 report = classification_report(y_test, y_pred, target_names=iris.target_names) print(report)

这里我想多花点篇幅讲混淆矩阵,因为它比单个准确率信息量大得多。混淆矩阵是一个 3×3 的矩阵,行是真实类别,列是预测类别。对角线上的数字代表预测正确的数量,非对角线数字代表混淆情况。在 KNN 模型下,你大概率能看到 Setosa 全部预测正确,而 Versicolor 和 Virginica 可能有 1 到 2 条互相预测错,这正好呼应了之前可视化时看到的“花瓣特征部分重叠”现象。

classification_report 里的精确率(precision)表示预测为某类的样本中有多少是真的;召回率(recall)表示某类真实样本中有多少被找出来了;F1 值是二者的调和平均。在类别不平衡的数据里,光看准确率会被“多数类”带偏,所以要配合这三个指标一起看。鸢尾花数据集类别均衡,准确率就够用了,但习惯要养好。

4. 模型调参与效果对比:让模型从“能用”到“好用”

如果你的代码到这步已经跑通,恭喜,你已经迈出了机器学习的第一步。但项目还没结束,因为 KNN 里的 K 值、距离度量方式都是可以调节的。调参不是玄学,而是理解模型行为的过程。

4.1 K 值如何影响分类结果

K 是 KNN 中最核心的超参数。K 太小(比如 1),模型对噪声非常敏感,一个异常点就可能让预测翻车,这叫过拟合;K 太大(比如 50),模型会把离得很远、根本不同类的样本也拉进来投票,边界过于平滑,这叫欠拟合。所以我们需要找一个合适的中间值。

一个朴素但有效的办法是:让 K 从 1 到 20 都试一遍,每次计算在测试集上的准确率,然后画成折线图。代码如下:

import numpy as np k_values = range(1, 21) accuracies = [] for k in k_values: knn_tmp = KNeighborsClassifier(n_neighbors=k) knn_tmp.fit(X_train, y_train) y_tmp_pred = knn_tmp.predict(X_test) accuracies.append(accuracy_score(y_test, y_tmp_pred)) plt.figure(figsize=(8, 4)) plt.plot(k_values, accuracies, marker='o') plt.xlabel('K 值') plt.ylabel('测试集准确率') plt.title('不同 K 值下的准确率') plt.xticks(k_values) plt.grid(True) plt.show() # 找到最高准确率对应的 K best_k = k_values[np.argmax(accuracies)] print(f"最佳 K 值: {best_k}, 准确率: {max(accuracies):.2f}")

画出来之后,你会发现准确率可能在某个 K 值到达 1.00,随后在 0.93~1.00 之间小幅波动。这是因为 30 条测试样本,每错 1 条,准确率就掉约 3.3 个百分点。所以别看到一个小波动就紧张,数据集小的时候,微小波动不代表模型本质变差了。我在项目里通常还会做一次交叉验证,把单次划分的偶然性进一步消除。

4.2 交叉验证:更可靠的模型评估方法

单次划分训练集和测试集,结果受 random_state 影响。哪怕设置了固定种子,也只能保证“可复现”,不能保证“没运气成分”。交叉验证(Cross-Validation)的做法是把训练数据切成 K 份(比如 5 份),轮流拿 4 份训练、1 份验证,算 5 次结果取平均。这样每个样本都有机会当验证数据,评估结果更稳定。

Scikit-learn 里调用 cross_val_score 只需要几行:

from sklearn.model_selection import cross_val_score knn_cv = KNeighborsClassifier(n_neighbors=5) scores = cross_val_score(knn_cv, iris.data, iris.target, cv=5) print("交叉验证得分:", scores) print(f"平均得分: {scores.mean():.3f}, 标准差: {scores.std():.3f}")

注意这里我用的是整个数据集 iris.data 和 iris.target,没有再单独分测试集,因为交叉验证会自动完成数据划分。5 折交叉验证跑完会输出 5 个得分,比如 [0.9667, 1.0, 0.9333, 0.9667, 1.0],平均 0.973。这个平均得分比单次划分的准确率更有说服力。

4.3 把多个模型放在一起对比

机器学习领域没有免费的午餐,没有哪个算法能在所有数据集上通吃。我建议你入门阶段多跑几个模型,横向对比一下效果。这里用同一份划分好的训练集和测试集,快速测试逻辑回归(Logistic Regression)、决策树(Decision Tree)和支持向量机(SVM):

from sklearn.linear_model import LogisticRegression from sklearn.tree import DecisionTreeClassifier from sklearn.svm import SVC models = { "KNN": KNeighborsClassifier(n_neighbors=5), "Logistic Regression": LogisticRegression(max_iter=200), "Decision Tree": DecisionTreeClassifier(random_state=42), "SVM": SVC() } for name, model in models.items(): model.fit(X_train, y_train) y_model_pred = model.predict(X_test) acc = accuracy_score(y_test, y_model_pred) print(f"{name} 测试集准确率: {acc:.2f}")

在鸢尾花数据集上,这几个模型的准确率大概率都在 0.93 以上。逻辑回归和 SVM 往往能到 1.00,决策树偶尔会稍差一点。原因在于决策树容易过拟合小数据集,需要限制树的深度(比如 max_depth=3)才能发挥更好效果。这个对比实验就是为了让你直观感受:不同模型的归纳偏置不同,同一个数据集上的表现也不同。理解这一点,比背十个算法的推导公式更重要。

5. 常见问题与排查技巧实录

既然是进阶篇的实操内容,我把自己和身边朋友实际踩过的一些坑整理一下。有些问题看似很小,但不解决会让人非常崩溃。

5.1 常见报错速查表

下面这张表列了入门阶段最常遇到的几个问题,以及排查思路。

问题现象可能原因解决办法
运行 import sklearn 报 ModuleNotFoundErrorscikit-learn 未安装执行 pip install scikit-learn
绘图时中文标签显示为方块matplotlib 默认字体不含中文在绘图前设置中文字体,或用英文标签
混淆矩阵输出是一堆数字看不出规律没意识到矩阵行列顺序记住“行是真实,列是预测”,配合 classification_report 一起看
模型准确率每次运行都不同train_test_split 没设 random_state加 random_state=42 固定随机种子
逻辑回归报 ConvergenceWarning默认迭代次数不够设置 max_iter=200 或更大
预测时传错数据维度predict 传入了单个样本而不是二维数组用 X_test[i].reshape(1, -1) 包装成二维

5.2 新手最容易忽视的细节

先说绘图中文乱码问题,这是最普遍的。解决方案很土但有效,在代码开头加上下面两行,把默认字体切换成系统中文字体:

plt.rcParams['font.sans-serif'] = ['SimHei', 'Arial Unicode MS', 'Microsoft YaHei'] plt.rcParams['axes.unicode_minus'] = False

第一行设置中文字体,第二行解决负号显示成方块的问题。要是你的系统里没有这些字体,干脆直接不写中文字标签,改用英文,省心。

再说一个特别容易被忽略的问题:预测形状错误。假设你想对测试集里的第一条样本做预测,直接写 model.predict(X_test[0]),会报“Expected 2D array, got 1D array instead”的错。因为 Scikit-learn 的 predict 接口期望输入是二维数组(样本数, 特征数),单条样本必须用 reshape(1, -1) 转成二维才能传进去。这个报错信息新手几乎都会遇到,看到之后心里不慌就行。

5.3 关于模型准确率的理性认识

在鸢尾花数据集上把准确率做到 95% 以上很容易,但你别急着高兴。小数据集上所有模型都可能有不错的表现,真正的考验是模型能不能泛化到新数据。我在实际项目中发现,很多人拿 Kaggle 经典数据集练手练得很爽,一上真实业务数据就翻车,原因多半是真实数据的分布更复杂、噪声更大、特征之间关联更隐蔽。

所以,每次拿到一个数据集,第一件事是观察样本量、特征类型、缺失情况、类别分布,第二件事是可视化,第三件事才是选模型。不要一上来就贪心用 XGBoost、神经网络,先把简单模型的基线跑出来,后续模型效果如果明显优于基线,才算有改进。

6. 个人体验与下一步扩展方向

鸢尾花分类跑通之后,说明你已经理解了机器学习的核心循环:加载数据、划分数据集、训练模型、评估性能、调参优化。剩下的就是在这里继续加减东西。我最推荐你做的第一个改动是:把数据集换掉。

Scikit-learn 里还有手写数字数据集(digits)、乳腺癌数据集(breast_cancer)等,换一个数据集,重新走一遍全流程。你会发现 KNN 在手写数字上的表现、调参过程、评估指标的关注点,都和鸢尾花不一样,因为数据维度变高了、类别变多了、数字图片还可能受噪声影响。我当年正是在换完数据集之后,才真正理解了特征工程的威力。

另外一个可以扩展的方向是数据标准化。拿鸢尾花数据练手时,标准化的提升不明显,但你可以故意把某个特征乘以 100,然后重新跑 KNN,再看看标准化前后的差距。自己做一次这种“破坏性实验”,比听别人说十遍“KNN 需要标准化”都管用。

最后说一句我的个人经验:不要迷信高准确率。机器学习的核心不是考试拿满分,而是模型在真实场景中稳定可靠地工作。我在实际项目中见过太多模型在测试集上得分漂亮、上线后却一塌糊涂的案例,根子往往就是数据划分不合理、评估方式不全面、对业务场景理解不足。从鸢尾花这个最小的项目开始,养成每一步都问“为什么这样做”的习惯,后面的路会顺畅得多。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/9 15:46:56

农产品追溯系统实战:从批次管理到二维码溯源全解析

简介:这套农产品追溯系统项目压缩包约3.95MB,包含1078个文件,主体为Java源码及编译后的class文件,配合JSP动态页面、JS/CSS前端资源与GIF/PNG/JPG图标图片,以及少量jar依赖和Eclipse工程配置,构成一套可运行…

作者头像 李华
网站建设 2026/9/9 15:46:45

开源机器鸭项目全解析:具身智能入门与仿真部署实战

最近社区里不少人在讨论 HuggingFace 最新开源的机器鸭项目:一只小鸭子模型,通过具身智能技术,在虚拟公寓里跑跳转圈,动作自然得像真鸭子撒欢。很多人第一反应是“玩具”,但认真看一遍技术栈就会发现,这只鸭…

作者头像 李华
网站建设 2026/9/9 15:46:36

Maven从入门到实战:依赖管理、构建生命周期与常见问题排查

如果你最近刚开始写 Java 项目,或者正被 IDEA 里一个叫Resolving Maven dependencies的进度条卡到怀疑人生,那这篇就是给你准备的。Maven 可以说是 Java 生态里最绕不开的基础设施了,它管你的 jar 包、管编译、管打包,甚至管整个项…

作者头像 李华
网站建设 2026/9/9 15:45:30

Java流程控制详解:分支循环与跳转实战

在Java的所有语法里,流程控制是我建议每个初学者第一个彻底吃透的知识点。它不像面向对象那样需要反复理解抽象概念,也不像集合框架那样需要背大量API,但几乎所有代码的执行逻辑都离不开它。哪怕你以后写的是业务代码、算法题,甚至…

作者头像 李华
网站建设 2026/9/9 15:44:58

LeetCode 216组合总和III:回溯算法剪枝与去重详解

刷题刷到第156天,题目编号是216。今天这题是回溯专题里的组合总和III,问题描述短得像散文:从数字1到9里选k个数,每个数字最多用一次,让这些数的和等于n,返回所有可能的组合。我一开始以为这就是组合总和的换…

作者头像 李华
网站建设 2026/9/9 15:43:48

STM32F103+HAL库模拟I2C驱动0.96寸OLED(SSD1306)教程

简介:面向嵌入式初学者的STM32F103C8T6开发例程,演示如何基于HAL库与模拟I2C驱动0.96英寸OLED显示屏。资源聚焦GPIO引脚模拟I2C时序、OLED初始化序列及字符/图形显示,适用于硬件I2C被占用或需灵活调整引脚的场景,适合正在学习STM3…

作者头像 李华