ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

入门实践工程一:基于 sklearn 的鸢尾花分类(传统机器学习入门)|附:环境依赖及工程源码

入门实践工程一:基于 sklearn 的鸢尾花分类(传统机器学习入门)|附:环境依赖及工程源码

🔥别再死磕枯燥理论了!AI 时代,拿实战作品说话才是硬道理!(原创不易哈,希望可以帮到还有些许学习劲儿的同学们)🔥
【进阶版还在创作中,耗费精力中……】

跳转到专栏目录,你学习更有方向和思路……

入门实践工程一:基于 sklearn 的鸢尾花分类(传统机器学习入门)|附:环境依赖及工程源码

简介

使用 scikit-learn 内置鸢尾花数据集,训练并对比 KNN / SVM / 随机森林三种经典分类器,评估准确率并绘制二维决策边界可视化图。是最轻量的 AI 入门项目,帮助理解「特征→模型→评估→可视化」传统机器学习全流程,与深度学习项目形成互补。

项目亮点:

  • 🚀零基础友好:无需深度学习框架,无需 GPU,安装即运行
  • 📊可视化直观:决策边界图直观展示不同模型的分类逻辑
  • ⚖️模型对比:三种经典算法横向对比,理解各自优缺点
  • 🎯完整闭环:从数据加载到模型评估,体验完整机器学习流程

鸢尾花数据集简介

鸢尾花(Iris)数据集是机器学习领域最经典的数据集之一,由统计学家 R.A. Fisher 在 1936 年引入。该数据集包含 150 个样本,每个样本有 4 个特征:

  1. 花萼长度(sepal length,单位:厘米)
  2. 花萼宽度(sepal width,单位:厘米)
  3. 花瓣长度(petal length,单位:厘米)
  4. 花瓣宽度(petal width,单位:厘米)

三个类别分别为:

  • Setosa(山鸢尾)
  • Versicolor(杂色鸢尾)
  • Virginica(维吉尼亚鸢尾)

该数据集的特点是类别间线性可分性良好,非常适合作为分类算法的入门实践。

工程详细介绍

核心思想

传统监督学习的完整闭环——用少量表格特征,借助「距离投票 / 最大间隔 / 集成」三类思想完成多分类,无需神经网络与 GPU,是理解「特征→拟合→评估→可视化」的最简载体,与深度学习项目互为补充。

实现方法

1. 数据准备
  • 数据源:sklearn 内置鸢尾花数据集(150 样本,4 维特征:花萼/花瓣的长宽,3 个类别)
  • 数据划分:采用留出法(Hold-out),按 7:3 比例划分训练集和测试集
  • 分层抽样:使用stratify=y确保每个类别的样本比例在划分后保持一致
2. 模型选择与对比

本项目对比三种经典分类算法,代表三种不同的分类思想:

K-最近邻(KNN)

  • 核心思想:“物以类聚” - 根据最近的 k 个邻居的类别进行投票
  • 优点:简单直观,无需训练过程
  • 缺点:预测时计算量大,对特征尺度敏感
  • 参数:k=5(经验值)

支持向量机(SVM)

  • 核心思想:寻找最大化类别间隔的超平面
  • 核函数:RBF(径向基函数)核,适合非线性分类
  • 优点:在高维空间表现优秀,泛化能力强
  • 参数:C=1.0(正则化参数),gamma=‘scale’

随机森林(Random Forest)

  • 核心思想:集成学习 - 多棵决策树投票决定最终结果
  • 优点:抗过拟合能力强,能处理高维特征
  • 缺点:模型可解释性较差
  • 参数:n_estimators=100(树的数量)
3. 训练与评估流程
  1. 数据加载与划分:加载数据集并按 7:3 划分
  2. 模型训练:分别用训练集训练三个模型
  3. 性能评估:在测试集上计算准确率
  4. 可视化分析:使用前两个特征绘制决策边界
4. 输出结果
  • 三种模型在测试集上的准确率对比
  • 决策边界可视化图(iris_decision_boundary.png
  • 一个示例预测,展示模型的实际应用

项目结构

01_iris_ml/ ├── main.py # 训练 + 对比 + 决策边界出图 ├── requirements.txt # 依赖包列表 └── iris_decision_boundary.png # 生成的决策边界图

环境配置与安装

系统要求

  • Python 3.7+
  • 任意操作系统(Windows/macOS/Linux)

安装依赖

pipinstallscikit-learn matplotlib numpy

注意事项:

  • 数据集由 sklearn 内置,无需联网下载,安装后即可运行
  • 所有依赖包均可通过 pip 一键安装
  • 无需 GPU 支持,普通 CPU 即可秒级完成训练

验证安装

importsklearnprint(f"scikit-learn 版本:{sklearn.__version__}")# 应该输出类似: scikit-learn 版本: 1.3.0

运行方式

方法一:直接运行(推荐)

python main.py

方法二:使用 requirements.txt

pipinstall-rrequirements.txt python main.py

运行过程解析

程序执行时会依次完成以下步骤:

  1. 数据加载:加载鸢尾花数据集并显示基本信息
  2. 数据划分:按 7:3 划分训练集和测试集
  3. 模型训练:依次训练 KNN、SVM、随机森林
  4. 性能评估:输出各模型在测试集上的准确率
  5. 可视化:生成决策边界对比图
  6. 示例预测:用最佳模型进行一个样本预测

代码详解

1. 数据加载与探索

iris=load_iris()X,y=iris.data,iris.target feature_names=iris.feature_names target_names=iris.target_names
  • X:特征矩阵,形状为 (150, 4)
  • y:标签向量,取值为 0、1、2
  • feature_names:四个特征的名称
  • target_names:三个类别的名称

2. 数据划分策略

X_train,X_test,y_train,y_test=train_test_split(X,y,test_size=0.3,random_state=42,stratify=y)
  • test_size=0.3:30% 数据作为测试集
  • random_state=42:固定随机种子,确保结果可复现
  • stratify=y:分层抽样,保持类别比例

3. 模型定义与训练

models={"KNN (k=5)":KNeighborsClassifier(n_neighbors=5),"SVM (RBF)":SVC(kernel="rbf",gamma="scale",C=1.0,probability=True),"随机森林 (100 棵树)":RandomForestClassifier(n_estimators=100,random_state=42),}

每个模型都有其独特的参数设置,这些参数基于经验值和数据集特性选择。

4. 决策边界可视化原理

# 创建网格点xx,yy=np.meshgrid(np.linspace(x_min,x_max,300),np.linspace(y_min,y_max,300))# 预测网格上每个点的类别Z=clf.predict(np.c_[xx.ravel(),yy.ravel()]).reshape(xx.shape)# 绘制等高线填充图ax.contourf(xx,yy,Z,alpha=0.3,cmap=plt.cm.Set1)

决策边界图通过以下步骤生成:

  1. 在特征空间创建密集的网格点
  2. 用训练好的模型预测每个网格点的类别
  3. 用不同颜色填充不同类别的区域
  4. 在图上叠加真实的测试样本点

预期结果

1. 控制台输出

运行程序后,控制台会显示类似以下信息:

鸢尾花数据集: 150 个样本, 4 个特征 特征: ['sepal length (cm)', 'sepal width (cm)', 'petal length (cm)', 'petal width (cm)'] 类别: ['setosa', 'versicolor', 'virginica'] KNN (k=5) 测试准确率 = 0.9778 SVM (RBF) 测试准确率 = 0.9778 随机森林 (100 棵树) 测试准确率 = 0.9556 决策边界对比图已保存到: /path/to/iris_decision_boundary.png 示例预测(KNN (k=5)): 样本=[[5.1, 3.5, 1.4, 0.2]] -> setosa

2. 生成的可视化图

程序会生成iris_decision_boundary.png文件,包含三张子图:

  • 左侧:KNN 决策边界(通常呈现不规则的区域划分)
  • 中间:SVM 决策边界(边界平滑,基于最大间隔原则)
  • 右侧:随机森林决策边界(可能呈现复杂的多区域划分)

每张图都显示:

  • 不同颜色的区域代表不同的预测类别
  • 散点代表测试集中的真实样本
  • 标题包含模型名称和仅使用前两个特征的测试准确率

结果分析与讨论

1. 准确率分析

鸢尾花数据集相对简单,三种模型通常都能达到 95% 以上的准确率:

  • KNN 和 SVM:在这个数据集上表现非常接近,经常达到 97-98% 的准确率
  • 随机森林:可能略低一些,但仍在 95% 以上

为什么准确率这么高?

  1. 数据集本身线性可分性良好
  2. 特征数量少(4个),样本数量适中(150个)
  3. 类别间差异明显,特别是 Setosa 与其他两类容易区分

2. 决策边界对比

通过决策边界图可以直观看到不同模型的分类逻辑:

KNN 决策边界特点:

  • 边界不规则,呈现"锯齿状"
  • 每个点的类别由其最近邻居决定
  • 对局部噪声敏感

SVM 决策边界特点:

  • 边界平滑,基于最大间隔原则
  • 使用 RBF 核可以处理非线性关系
  • 泛化能力较强

随机森林决策边界特点:

  • 可能呈现多个小区域
  • 基于多棵树的投票结果
  • 对异常值相对鲁棒

3. 模型选择建议

对于鸢尾花分类任务:

  • 追求简单快速:选择 KNN,无需调参,实现简单
  • 追求泛化能力:选择 SVM,特别是面对新数据时
  • 追求稳定鲁棒:选择随机森林,对噪声和异常值不敏感

常见问题与解决方案

Q1: 安装 scikit-learn 失败

解决方案:

# 使用国内镜像源pipinstallscikit-learn matplotlib numpy-ihttps://pypi.tuna.tsinghua.edu.cn/simple# 或使用 condacondainstallscikit-learn matplotlib numpy

Q2: 运行时报错 “ModuleNotFoundError”

可能原因:依赖包未正确安装
解决方案:

# 检查已安装的包pip list|grep-E"scikit-learn|matplotlib|numpy"# 重新安装pip uninstall scikit-learn matplotlib numpy pipinstallscikit-learn matplotlib numpy

Q3: 生成的图片无法显示或保存

解决方案:

# 在代码开头添加以下配置importmatplotlib matplotlib.use('Agg')# 使用非交互式后端

Q4: 准确率每次运行都不一样

原因:未设置随机种子
解决方案:代码中已设置random_state=42,确保结果可复现

扩展方向与进阶学习

1. 特征工程扩展

# 添加特征标准化fromsklearn.preprocessingimportStandardScaler scaler=StandardScaler()X_scaled=scaler.fit_transform(X)# 添加 PCA 降维可视化fromsklearn.decompositionimportPCA pca=PCA(n_components=2)X_pca=pca.fit_transform(X)

2. 更换数据集挑战

  • 葡萄酒数据集:13个特征,3个类别,特征间相关性更强
  • 手写数字数据集:64个特征(8×8像素),10个类别,更适合复杂模型
  • 乳腺癌数据集:30个特征,二分类问题,适合逻辑回归等算法

3. 模型扩展与对比

# 添加逻辑回归fromsklearn.linear_modelimportLogisticRegression models["逻辑回归"]=LogisticRegression(max_iter=1000)# 添加 XGBoostfromxgboostimportXGBClassifier models["XGBoost"]=XGBClassifier(n_estimators=100)# 绘制 ROC 曲线(二分类)fromsklearn.metricsimportroc_curve,auc fpr,tpr,_=roc_curve(y_test_binary,y_score)roc_auc=auc(fpr,tpr)

4. 交叉验证与超参数调优

fromsklearn.model_selectionimportcross_val_score,GridSearchCV# K 折交叉验证scores=cross_val_score(model,X,y,cv=5)# 网格搜索调参param_grid={'n_neighbors':[3,5,7,9]}grid_search=GridSearchCV(KNeighborsClassifier(),param_grid,cv=5)grid_search.fit(X_train,y_train)

5. 模型可解释性

# 随机森林特征重要性importances=rf_model.feature_importances_ indices=np.argsort(importances)[::-1]# 绘制特征重要性图plt.figure()plt.title("特征重要性")plt.bar(range(X.shape[1]),importances[indices])plt.xticks(range(X.shape[1]),[feature_names[i]foriinindices],rotation=45)plt.tight_layout()

学习建议与下一步

给初学者的建议

  1. 先运行再理解:不要被代码吓到,先运行起来看到结果
  2. 逐行调试:在关键位置添加print()语句,查看中间结果
  3. 修改参数:尝试修改 k 值、树的数量等参数,观察结果变化
  4. 可视化探索:使用 matplotlib 绘制更多图表,如特征分布、混淆矩阵等

知识体系构建

完成本项目后,建议按以下路径继续学习:

基础巩固(1-2周)

  1. 理解监督学习的基本概念:特征、标签、训练、测试
  2. 掌握数据预处理:缺失值处理、特征缩放、编码
  3. 学习模型评估指标:准确率、精确率、召回率、F1 分数

技能提升(2-4周)

  1. 尝试其他分类算法:朴素贝叶斯、决策树、梯度提升
  2. 学习回归问题:线性回归、多项式回归
  3. 了解聚类算法:K-means、DBSCAN

项目实践(1-2个月)

  1. 参加 Kaggle 入门竞赛(如 Titanic、House Prices)
  2. 尝试真实业务数据(如用户流失预测、信用评分)
  3. 学习模型部署:使用 Flask/FastAPI 部署简单模型

工程源码

main.py

""" 入门实践工程一:基于 sklearn 的鸢尾花分类(传统机器学习入门) ========================================================= 使用 scikit-learn 内置的鸢尾花数据集,训练并对比 KNN / SVM / 随机森林 三种经典分类器,评估准确率,并绘制二维决策边界可视化图。 全程无需深度学习框架、无需联网,是最轻量的 AI 入门项目。 运行: python main.py """importosimportmatplotlib matplotlib.use("Agg")importmatplotlib.pyplotaspltimportnumpyasnpfromsklearn.datasetsimportload_irisfromsklearn.ensembleimportRandomForestClassifierfromsklearn.model_selectionimporttrain_test_splitfromsklearn.neighborsimportKNeighborsClassifierfromsklearn.svmimportSVC BASE_DIR=os.path.dirname(os.path.abspath(__file__))defmain():iris=load_iris()X,y=iris.data,iris.target feature_names=iris.feature_names target_names=iris.target_namesprint(f"鸢尾花数据集:{X.shape[0]}个样本,{X.shape[1]}个特征")print(f"特征:{feature_names}")print(f"类别:{list(target_names)}\n")X_train,X_test,y_train,y_test=train_test_split(X,y,test_size=0.3,random_state=42,stratify=y)models={"KNN (k=5)":KNeighborsClassifier(n_neighbors=5),"SVM (RBF)":SVC(kernel="rbf",gamma="scale",C=1.0,probability=True),"随机森林 (100 棵树)":RandomForestClassifier(n_estimators=100,random_state=42),}results={}forname,clfinmodels.items():clf.fit(X_train,y_train)acc=clf.score(X_test,y_test)results[name]=accprint(f"{name:<22}测试准确率 ={acc:.4f}")# ---- 决策边界可视化(取前两个特征,便于 2D 绘图) ----X2=X[:,:2]# 花萼长度 + 花萼宽度X_tr2,X_te2,y_tr2,y_te2=train_test_split(X2,y,test_size=0.3,random_state=42,stratify=y)fig,axes=plt.subplots(1,len(models),figsize=(16,5))x_min,x_max=X2[:,0].min()-0.5,X2[:,0].max()+0.5y_min,y_max=X2[:,1].min()-0.5,X2[:,1].max()+0.5xx,yy=np.meshgrid(np.linspace(x_min,x_max,300),np.linspace(y_min,y_max,300))forax,(name,_)inzip(axes,models.items()):clf=models[name]clf.fit(X_tr2,y_tr2)Z=clf.predict(np.c_[xx.ravel(),yy.ravel()]).reshape(xx.shape)ax.contourf(xx,yy,Z,alpha=0.3,cmap=plt.cm.Set1)scatter=ax.scatter(X_te2[:,0],X_te2[:,1],c=y_te2,cmap=plt.cm.Set1,edgecolors="k",s=40)acc2=clf.score(X_te2,y_te2)ax.set_title(f"{name}\n(2特征 测试准确率={acc2:.3f})")ax.set_xlabel(feature_names[0])ax.set_ylabel(feature_names[1])fig.suptitle("鸢尾花分类决策边界对比(仅用前两个特征)",fontsize=14)plt.tight_layout()fig_path=os.path.join(BASE_DIR,"iris_decision_boundary.png")plt.savefig(fig_path,dpi=120)print(f"\n决策边界对比图已保存到:{fig_path}")# ---- 用全特征模型做一个示例预测 ----best_name=max(results,key=results.get)best_model=models[best_name]sample=np.array([[5.1,3.5,1.4,0.2]])# 典型山鸢尾pred=best_model.predict(sample)[0]print(f"\n示例预测({best_name}): 样本={sample.tolist()}->{target_names[pred]}")if__name__=="__main__":main()
返回列表