【机器学习】利用KNN算法进行鸢尾花分类
1、鸢尾花分类问题介绍
鸢尾花分类是机器学习领域入门级经典多分类任务,由统计学家罗纳德・费歇尔(Ronald Fisher)于 1936 年提出,堪称分类算法的「Hello World」。其核心问题是:通过鸢尾花可量化的 4 个物理特征,实现对 3 个鸢尾花品种的自动化、精准分类。
在园艺识别、植物育种等实际场景中,人工区分鸢尾花品种需专业植物学知识,且易受主观判断、植株生长状态影响;而通过机器学习算法构建分类模型,可将物理特征转化为客观分类依据,实现低成本、高效率的品种识别,这也是该问题的实际应用价值。

本次实战选择 KNN 算法解决该问题,是算法特性与问题场景的高度适配:
- 问题为「小样本 + 低维度」(150 条样本、4 个特征),完美契合 KNN 的适用范围,规避其大数据集下的效率短板;
- 3 类鸢尾花特征区分度高,无需复杂特征工程,KNN 的「距离度量 + 多数表决」逻辑即可实现高准确率,新手易出效果;
- KNN 是惰性学习算法,无需复杂训练过程,代码实现简洁,能让初学者聚焦分类任务的完整流程;
- 数据集无缺失值、无严重异常值,数据质量高,减少预处理工作量,重点掌握 KNN 核心操作(标准化、参数选择)
二、数据集核心解析
鸢尾花(Iris)数据集是 Sklearn 内置的标准测试集,无需额外下载,核心信息如下,是入门分类任务的理想数据集:
- 样本规模:共 150 条样本,3 个鸢尾花品种各 50 条,类别分布均匀,无样本不平衡问题;
- 特征维度:4 个连续型数值特征(单位:cm),均为可直接测量的物理特征,无冗余维度:
- sepal length(花萼长度)
- sepal width(花萼宽度)
- petal length(花瓣长度)
- petal width(花瓣宽度)
- 标签类别:3 个离散型标签,对应 3 种鸢尾花品种,标签与品种一一映射:0 → Iris-setosa(山鸢尾)1 → Iris-versicolor(变色鸢尾)2 → Iris-virginica(维吉尼亚鸢尾)
- 数据质量:无缺失值、无异常值、无重复样本,可直接用于模型训练,降低新手操作门槛。
鸢尾花数据集(部分数据效果如下)介绍 
三、完整实现代码(零基础可运行)
本次实战实现从数据加载到模型部署的全流程,代码注释详细,依赖仅需sklearn、numpy、matplotlib,直接复制即可运行,核心流程为:加载数据→划分数据集→特征标准化→训练 KNN→模型评估→新样本预测。
3.1 完整代码
# 1. 导入必备工具包
import numpy as np
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.neighbors import KNeighborsClassifier
from sklearn.metrics import accuracy_score, classification_report, confusion_matrix
# 2. 加载并查看鸢尾花数据集
iris = load_iris()
X = iris.data # 特征数据:150行4列,对应4个物理特征
y = iris.target # 标签数据:150行1列,对应3个品种
# 打印数据集基础信息,快速了解数据结构
print("===== 鸢尾花数据集基础信息 =====")
print(f"特征数据形状:{X.shape} | 标签数据形状:{y.shape}")
print(f"特征名称:{iris.feature_names}")
print(f"品种名称:{iris.target_names}")
print(f"前3条样本示例:\n特征:{X[:3]} \n标签:{y[:3]} \n对应品种:{iris.target_names[y[:3]]}")
# 3. 划分训练集与测试集(核心参数保证数据合理性)
# test_size=0.2:测试集占20%(30条),训练集占80%(120条)
# random_state=22:固定随机种子,保证结果可复现,新手调试必备
# stratify=y:按标签比例划分,避免训练/测试集类别分布不均
X_train, X_test, y_train, y_test = train_test_split(
X, y, test_size=0.2, random_state=22, stratify=y
)
print(f"\n===== 数据集划分结果 =====")
print(f"训练集:特征{X_train.shape} | 标签{y_train.shape}")
print(f"测试集:特征{X_test.shape} | 标签{y_test.shape}")
# 4. 特征标准化(KNN算法必须步骤,核心:消除量纲影响)
# 注意:测试集仅做转换,不重新拟合,避免数据泄露
scaler = StandardScaler()
X_train_scaled = scaler.fit_transform(X_train) # 训练集:拟合+转换
X_test_scaled = scaler.transform(X_test) # 测试集:仅转换
print(f"\n===== 标准化后数据特征 =====")
print(f"训练集各特征均值:{np.round(scaler.mean_, 2)}(趋近0)")
print(f"训练集各特征标准差:{np.round(scaler.scale_, 2)}(趋近1)")
# 5. 初始化并训练KNN分类模型(自定义核心参数,贴合实战)
# 结合前文理论,选择距离加权投票,提升预测精准度
knn_model = KNeighborsClassifier(
n_neighbors=5, # 核心超参数:选取5个近邻
weights='distance', # 权重方式:距离越近,投票权重越高
algorithm='auto', # 近邻搜索算法:自动选择最优(推荐)
p=2, # 距离p值:p=2对应欧式距离(KNN默认)
metric='euclidean' # 距离度量:显式指定欧式距离,更直观
)
# KNN是惰性学习,fit方法仅存储训练数据,无实际训练过程
knn_model.fit(X_train_scaled, y_train)
print(f"\n===== KNN模型训练完成 =====")
print(f"模型训练使用特征:{iris.feature_names}")
# 6. 模型预测与多维度评估(全面验证模型效果)
y_pred = knn_model.predict(X_test_scaled) # 测试集预测结果
# 6.1 计算整体准确率
acc = accuracy_score(y_test, y_pred)
# 6.2 生成详细分类报告(精确率、召回率、F1值)
class_report = classification_report(y_test, y_pred, target_names=iris.target_names)
# 6.3 生成混淆矩阵(直观查看类别错分情况)
cm = confusion_matrix(y_test, y_pred)
# 打印评估结果
print(f"\n===== 模型测试集评估结果 =====")
print(f"模型整体准确率:{acc:.4f}(1.0表示100%正确)")
print(f"\n详细分类报告(精确率/召回率/F1值):\n{class_report}")
print(f"混淆矩阵(行=真实类别,列=预测类别):\n{cm}")
# 7. 混淆矩阵可视化(更直观理解分类结果,可选)
plt.figure(figsize=(8, 6), dpi=100)
plt.imshow(cm, interpolation='nearest', cmap=plt.cm.Blues)
plt.title('KNN鸢尾花分类 - 混淆矩阵', fontsize=14)
plt.colorbar() # 颜色刻度条
# 设置坐标轴标签为品种名称
tick_marks = np.arange(len(iris.target_names))
plt.xticks(tick_marks, iris.target_names, rotation=45, fontsize=12)
plt.yticks(tick_marks, iris.target_names, fontsize=12)
# 混淆矩阵内标注数值
thresh = cm.max() / 2.0
for i in range(cm.shape[0]):
for j in range(cm.shape[1]):
plt.text(j, i, format(cm[i, j], 'd'),
ha="center", va="center", fontsize=12,
color="white" if cm[i, j] > thresh else "black")
plt.ylabel('真实品种', fontsize=12)
plt.xlabel('预测品种', fontsize=12)
plt.tight_layout() # 适配布局,防止标签遮挡
plt.show()
# 8. 新样本预测(模拟实际应用场景,模型部署核心)
print(f"\n===== 新样本预测演示(模拟实际应用) =====")
# 模拟1条新的鸢尾花特征数据(花萼长、花萼宽、花瓣长、花瓣宽)
new_iris = np.array([[5.1, 3.5, 1.4, 0.2]]) # 典型山鸢尾特征
# 注意:新样本必须与训练集做相同的标准化处理
new_iris_scaled = scaler.transform(new_iris)
# 预测品种标签和各类别概率
pred_label = knn_model.predict(new_iris_scaled)
pred_proba = knn_model.predict_proba(new_iris_scaled) # 预测各类别概率
# 打印预测结果
print(f"新样本特征:{new_iris[0]}(单位:cm)")
print(f"预测品种:{iris.target_names[pred_label[0]]}")
print(f"各品种预测概率:{dict(zip(iris.target_names, np.round(pred_proba[0], 4)))}")
3.2 环境依赖安装
若本地缺少相关库,执行以下命令快速安装(清华镜像源,速度更快):
pip install numpy matplotlib scikit-learn -i https://pypi.tuna.tsinghua.edu.cn/simple/
四、代码运行结果与核心解读
4.1 控制台输出结果
===== 鸢尾花数据集基础信息 =====
特征数据形状:(150, 4) | 标签数据形状:(150,)
特征名称:['sepal length (cm)', 'sepal width (cm)', 'petal length (cm)', 'petal width (cm)']
品种名称:['setosa' 'versicolor' 'virginica']
前3条样本示例:
特征:[[5.1 3.5 1.4 0.2]
[4.9 3. 1.4 0.2]
[4.7 3.2 1.3 0.2]]
标签:[0 0 0]
对应品种:['setosa' 'setosa' 'setosa']
===== 数据集划分结果 =====
训练集:特征(120, 4) | 标签(120,)
测试集:特征(30, 4) | 标签(30,)
===== 标准化后数据特征 =====
训练集各特征均值:[5.89 3.05 3.78 1.21](趋近0)
训练集各特征标准差:[0.83 0.43 1.75 0.76](趋近1)
===== KNN模型训练完成 =====
模型训练使用特征:['sepal length (cm)', 'sepal width (cm)', 'petal length (cm)', 'petal width (cm)']
===== 模型测试集评估结果 =====
模型整体准确率:1.0000(1.0表示100%正确)
详细分类报告(精确率/召回率/F1值):
precision recall f1-score support
setosa 1.00 1.00 1.00 10
versicolor 1.00 1.00 1.00 10
virginica 1.00 1.00 1.00 10
accuracy 1.00 30
macro avg 1.00 1.00 1.00 30
weighted avg 1.00 1.00 1.00 30
混淆矩阵(行=真实类别,列=预测类别):
[[10 0 0]
[ 0 10 0]
[ 0 0 10]]
===== 新样本预测演示(模拟实际应用) =====
新样本特征:[5.1 3.5 1.4 0.2](单位:cm)
预测品种:setosa
各品种预测概率:{'setosa': 1.0, 'versicolor': 0.0, 'virginica': 0.0}
4.2 核心结果解读
- 数据集与划分:150 条样本按 8:2 划分为训练集(120 条)和测试集(30 条),
stratify=y保证每类品种在训练 / 测试集中各 10 条,分布均匀; - 标准化效果:标准化后训练集特征均值趋近 0、标准差趋近 1,成功消除了 4 个特征的量纲差异,避免距离计算被数值大的特征主导;
- 模型准确率:测试集准确率 100%(1.0000),这是因为鸢尾花数据集特征区分度极高,是入门算法的理想测试集(实际业务场景中 95%+ 即为优秀);
- 详细分类报告:3 个品种的精确率、召回率、F1 值均为 1.0,说明模型对每一类品种的分类效果都完美,无错分情况;
- 混淆矩阵:对角线数值为 10,其余为 0,代表所有样本都被正确分类,无跨品种错分,直观验证了模型的准确性;
- 新样本预测:模拟的山鸢尾样本被精准识别,且预测概率为 1.0,说明模型对典型样本的判断具有 100% 的置信度。
4.3 混淆矩阵可视化结果
运行代码后会弹出可视化窗口,横轴为「预测品种」,纵轴为「真实品种」,颜色越深代表样本数量越多,对角线全为深蓝色且数值为 10,代表无错分,可视化结果比纯数字更直观,便于向非技术人员展示模型效果。
五、关键优化:网格搜索 + 交叉验证选最优参数
上述代码使用经验参数n_neighbors=5已实现完美效果,但在实际业务场景中,参数选择直接影响模型性能,因此需要通过网格搜索遍历参数组合,结合交叉验证选择最优参数,这是机器学习调优的通用方法。
5.1 参数优化代码
# 导入网格搜索与交叉验证工具
from sklearn.model_selection import GridSearchCV
# 1. 定义待搜索的参数范围
# 遍历不同K值、权重方式、距离类型,找到最优组合
param_grid = {
'n_neighbors': [3, 5, 7, 9, 11], # 测试不同近邻数
'weights': ['uniform', 'distance'], # 测试等权重/距离加权
'p': [1, 2] # 测试曼哈顿距离(p=1)/欧式距离(p=2)
}
# 2. 初始化网格搜索对象
# cv=5:5折交叉验证,将训练集分成5份,4份训练1份验证,循环5次
# scoring='accuracy':以准确率为评价指标
# n_jobs=-1:利用所有CPU核心加速,适合大数据集
grid_search = GridSearchCV(
estimator=KNeighborsClassifier(algorithm='auto'),
param_grid=param_grid,
cv=5,
scoring='accuracy',
n_jobs=-1,
verbose=1 # 打印搜索过程,便于查看进度
)
# 3. 在训练集上执行网格搜索
grid_search.fit(X_train_scaled, y_train)
# 4. 输出最优结果
print("===== 网格搜索+交叉验证 - 最优参数结果 =====")
print(f"最优参数组合:{grid_search.best_params_}")
print(f"交叉验证最优准确率:{grid_search.best_score_:.4f}")
print(f"最优模型测试集准确率:{grid_search.best_estimator_.score(X_test_scaled, y_test):.4f}")
# 5. 使用最优模型进行新样本预测
best_knn = grid_search.best_estimator_
new_pred = best_knn.predict(new_iris_scaled)
print(f"\n最优模型新样本预测结果:{iris.target_names[new_pred[0]]}")
5.2 优化结果输出
Fitting 5 folds for each of 20 candidates, totalling 100 fits
===== 网格搜索+交叉验证 - 最优参数结果 =====
最优参数组合:{'n_neighbors': 5, 'p': 2, 'weights': 'distance'}
交叉验证最优准确率:0.9917
最优模型测试集准确率:1.0000
最优模型新样本预测结果:setosa
5.3 优化结果解读
- 搜索规模:本次共搜索 20 个参数组合(5 个 K 值 ×2 种权重 ×2 种距离),5 折交叉验证共执行 100 次模型训练,全面覆盖主流参数;
- 最优参数:网格搜索最终选择
n_neighbors=5、p=2、weights='distance',与我们之前的经验参数一致,验证了经验参数的合理性; - 泛化能力:交叉验证最优准确率为 99.17%,测试集准确率为 100%,说明模型无过拟合,泛化能力优秀;
- 通用价值:该调优方法适用于所有机器学习算法,即使在复杂业务场景中,也能通过网格搜索找到最优参数,是新手必须掌握的调优技巧。
六、实战核心总结与关键注意点
6.1 实战核心总结
本次 KNN 实现鸢尾花分类的通用流程,适用于绝大多数分类任务,可牢记:
加载数据 → 划分训练/测试集(分层划分) → 特征标准化(KNN必须) → 初始化模型 → 训练模型 → 多维度评估 → 参数优化 → 新样本预测
6.2 关键注意点(新手避坑指南)
- 特征标准化是 KNN 的必做步骤:KNN 基于距离计算,若特征量纲不同(如花瓣长度 cm、花瓣宽度 mm),数值大的特征会主导距离计算,导致模型失效;
- 划分数据集必加
stratify=y:避免训练集 / 测试集类别分布不均(如训练集全是 A 类,测试集全是 B 类),导致模型训练无意义; - 固定
random_state:保证每次运行代码的结果一致,便于新手调试和问题排查,实际部署时可删除; - 测试集仅做标准化转换:测试集必须使用训练集的标准化参数(均值、标准差),否则会导致数据泄露,模型泛化能力下降;
- K 值选择需结合交叉验证:不要直接使用默认 K=5,实际业务场景中需通过网格搜索遍历 K 值,找到最适合当前数据集的参数;
- 权重优先选择
distance:距离加权投票比等权重投票(uniform)更精准,尤其是样本分布不均时,能让更近的样本发挥更大作用。
6.3 KNN 算法实战心得
KNN 作为入门级算法,虽原理简单,但包含了机器学习的核心思想:从数据中学习规律,实现对未知数据的预测。本次实战中,KNN 在鸢尾花数据集上的完美表现,印证了其在「小样本、低维度、类别分布均匀」场景中的优势。
同时也需认识到 KNN 的局限性:若面对大数据集(10 万 + 样本)、高维数据(100 + 特征),KNN 的预测效率会大幅下降,此时需选择决策树、随机森林、神经网络等更高效的算法,这也说明「算法无优劣,适配才是关键」。
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐



所有评论(0)