ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

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

入门实践工程一:基于 sklearn 的鸢尾花分类(传统机器学习入门)|附:环境依赖及工程源码 别再死磕枯燥理论了AI 时代拿实战作品说话才是硬道理原创不易哈希望可以帮到还有些许学习劲儿的同学们【进阶版还在创作中耗费精力中……】跳转到专栏目录你学习更有方向和思路……入门实践工程一基于 sklearn 的鸢尾花分类传统机器学习入门|附:环境依赖及工程源码简介使用 scikit-learn 内置鸢尾花数据集训练并对比 KNN / SVM / 随机森林三种经典分类器评估准确率并绘制二维决策边界可视化图。是最轻量的 AI 入门项目帮助理解「特征→模型→评估→可视化」传统机器学习全流程与深度学习项目形成互补。项目亮点零基础友好无需深度学习框架无需 GPU安装即运行可视化直观决策边界图直观展示不同模型的分类逻辑⚖️模型对比三种经典算法横向对比理解各自优缺点完整闭环从数据加载到模型评估体验完整机器学习流程鸢尾花数据集简介鸢尾花Iris数据集是机器学习领域最经典的数据集之一由统计学家 R.A. Fisher 在 1936 年引入。该数据集包含 150 个样本每个样本有 4 个特征花萼长度sepal length单位厘米花萼宽度sepal width单位厘米花瓣长度petal length单位厘米花瓣宽度petal width单位厘米三个类别分别为Setosa山鸢尾Versicolor杂色鸢尾Virginica维吉尼亚鸢尾该数据集的特点是类别间线性可分性良好非常适合作为分类算法的入门实践。工程详细介绍核心思想传统监督学习的完整闭环——用少量表格特征借助「距离投票 / 最大间隔 / 集成」三类思想完成多分类无需神经网络与 GPU是理解「特征→拟合→评估→可视化」的最简载体与深度学习项目互为补充。实现方法1. 数据准备数据源sklearn 内置鸢尾花数据集150 样本4 维特征花萼/花瓣的长宽3 个类别数据划分采用留出法Hold-out按 7:3 比例划分训练集和测试集分层抽样使用stratifyy确保每个类别的样本比例在划分后保持一致2. 模型选择与对比本项目对比三种经典分类算法代表三种不同的分类思想K-最近邻KNN核心思想“物以类聚” - 根据最近的 k 个邻居的类别进行投票优点简单直观无需训练过程缺点预测时计算量大对特征尺度敏感参数k5经验值支持向量机SVM核心思想寻找最大化类别间隔的超平面核函数RBF径向基函数核适合非线性分类优点在高维空间表现优秀泛化能力强参数C1.0正则化参数gamma‘scale’随机森林Random Forest核心思想集成学习 - 多棵决策树投票决定最终结果优点抗过拟合能力强能处理高维特征缺点模型可解释性较差参数n_estimators100树的数量3. 训练与评估流程数据加载与划分加载数据集并按 7:3 划分模型训练分别用训练集训练三个模型性能评估在测试集上计算准确率可视化分析使用前两个特征绘制决策边界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(fscikit-learn 版本:{sklearn.__version__})# 应该输出类似: scikit-learn 版本: 1.3.0运行方式方法一直接运行推荐python main.py方法二使用 requirements.txtpipinstall-rrequirements.txt python main.py运行过程解析程序执行时会依次完成以下步骤数据加载加载鸢尾花数据集并显示基本信息数据划分按 7:3 划分训练集和测试集模型训练依次训练 KNN、SVM、随机森林性能评估输出各模型在测试集上的准确率可视化生成决策边界对比图示例预测用最佳模型进行一个样本预测代码详解1. 数据加载与探索irisload_iris()X,yiris.data,iris.target feature_namesiris.feature_names target_namesiris.target_namesX特征矩阵形状为 (150, 4)y标签向量取值为 0、1、2feature_names四个特征的名称target_names三个类别的名称2. 数据划分策略X_train,X_test,y_train,y_testtrain_test_split(X,y,test_size0.3,random_state42,stratifyy)test_size0.330% 数据作为测试集random_state42固定随机种子确保结果可复现stratifyy分层抽样保持类别比例3. 模型定义与训练models{KNN (k5):KNeighborsClassifier(n_neighbors5),SVM (RBF):SVC(kernelrbf,gammascale,C1.0,probabilityTrue),随机森林 (100 棵树):RandomForestClassifier(n_estimators100,random_state42),}每个模型都有其独特的参数设置这些参数基于经验值和数据集特性选择。4. 决策边界可视化原理# 创建网格点xx,yynp.meshgrid(np.linspace(x_min,x_max,300),np.linspace(y_min,y_max,300))# 预测网格上每个点的类别Zclf.predict(np.c_[xx.ravel(),yy.ravel()]).reshape(xx.shape)# 绘制等高线填充图ax.contourf(xx,yy,Z,alpha0.3,cmapplt.cm.Set1)决策边界图通过以下步骤生成在特征空间创建密集的网格点用训练好的模型预测每个网格点的类别用不同颜色填充不同类别的区域在图上叠加真实的测试样本点预期结果1. 控制台输出运行程序后控制台会显示类似以下信息鸢尾花数据集: 150 个样本, 4 个特征 特征: [sepal length (cm), sepal width (cm), petal length (cm), petal width (cm)] 类别: [setosa, versicolor, virginica] KNN (k5) 测试准确率 0.9778 SVM (RBF) 测试准确率 0.9778 随机森林 (100 棵树) 测试准确率 0.9556 决策边界对比图已保存到: /path/to/iris_decision_boundary.png 示例预测KNN (k5): 样本[[5.1, 3.5, 1.4, 0.2]] - setosa2. 生成的可视化图程序会生成iris_decision_boundary.png文件包含三张子图左侧KNN 决策边界通常呈现不规则的区域划分中间SVM 决策边界边界平滑基于最大间隔原则右侧随机森林决策边界可能呈现复杂的多区域划分每张图都显示不同颜色的区域代表不同的预测类别散点代表测试集中的真实样本标题包含模型名称和仅使用前两个特征的测试准确率结果分析与讨论1. 准确率分析鸢尾花数据集相对简单三种模型通常都能达到 95% 以上的准确率KNN 和 SVM在这个数据集上表现非常接近经常达到 97-98% 的准确率随机森林可能略低一些但仍在 95% 以上为什么准确率这么高数据集本身线性可分性良好特征数量少4个样本数量适中150个类别间差异明显特别是 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 numpyQ2: 运行时报错 “ModuleNotFoundError”可能原因依赖包未正确安装解决方案# 检查已安装的包pip list|grep-Escikit-learn|matplotlib|numpy# 重新安装pip uninstall scikit-learn matplotlib numpy pipinstallscikit-learn matplotlib numpyQ3: 生成的图片无法显示或保存解决方案# 在代码开头添加以下配置importmatplotlib matplotlib.use(Agg)# 使用非交互式后端Q4: 准确率每次运行都不一样原因未设置随机种子解决方案代码中已设置random_state42确保结果可复现扩展方向与进阶学习1. 特征工程扩展# 添加特征标准化fromsklearn.preprocessingimportStandardScaler scalerStandardScaler()X_scaledscaler.fit_transform(X)# 添加 PCA 降维可视化fromsklearn.decompositionimportPCA pcaPCA(n_components2)X_pcapca.fit_transform(X)2. 更换数据集挑战葡萄酒数据集13个特征3个类别特征间相关性更强手写数字数据集64个特征8×8像素10个类别更适合复杂模型乳腺癌数据集30个特征二分类问题适合逻辑回归等算法3. 模型扩展与对比# 添加逻辑回归fromsklearn.linear_modelimportLogisticRegression models[逻辑回归]LogisticRegression(max_iter1000)# 添加 XGBoostfromxgboostimportXGBClassifier models[XGBoost]XGBClassifier(n_estimators100)# 绘制 ROC 曲线二分类fromsklearn.metricsimportroc_curve,auc fpr,tpr,_roc_curve(y_test_binary,y_score)roc_aucauc(fpr,tpr)4. 交叉验证与超参数调优fromsklearn.model_selectionimportcross_val_score,GridSearchCV# K 折交叉验证scorescross_val_score(model,X,y,cv5)# 网格搜索调参param_grid{n_neighbors:[3,5,7,9]}grid_searchGridSearchCV(KNeighborsClassifier(),param_grid,cv5)grid_search.fit(X_train,y_train)5. 模型可解释性# 随机森林特征重要性importancesrf_model.feature_importances_ indicesnp.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],rotation45)plt.tight_layout()学习建议与下一步给初学者的建议先运行再理解不要被代码吓到先运行起来看到结果逐行调试在关键位置添加print()语句查看中间结果修改参数尝试修改 k 值、树的数量等参数观察结果变化可视化探索使用 matplotlib 绘制更多图表如特征分布、混淆矩阵等知识体系构建完成本项目后建议按以下路径继续学习基础巩固1-2周理解监督学习的基本概念特征、标签、训练、测试掌握数据预处理缺失值处理、特征缩放、编码学习模型评估指标准确率、精确率、召回率、F1 分数技能提升2-4周尝试其他分类算法朴素贝叶斯、决策树、梯度提升学习回归问题线性回归、多项式回归了解聚类算法K-means、DBSCAN项目实践1-2个月参加 Kaggle 入门竞赛如 Titanic、House Prices尝试真实业务数据如用户流失预测、信用评分学习模型部署使用 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_DIRos.path.dirname(os.path.abspath(__file__))defmain():irisload_iris()X,yiris.data,iris.target feature_namesiris.feature_names target_namesiris.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_testtrain_test_split(X,y,test_size0.3,random_state42,stratifyy)models{KNN (k5):KNeighborsClassifier(n_neighbors5),SVM (RBF):SVC(kernelrbf,gammascale,C1.0,probabilityTrue),随机森林 (100 棵树):RandomForestClassifier(n_estimators100,random_state42),}results{}forname,clfinmodels.items():clf.fit(X_train,y_train)accclf.score(X_test,y_test)results[name]accprint(f{name:22}测试准确率 {acc:.4f})# ---- 决策边界可视化取前两个特征便于 2D 绘图 ----X2X[:,:2]# 花萼长度 花萼宽度X_tr2,X_te2,y_tr2,y_te2train_test_split(X2,y,test_size0.3,random_state42,stratifyy)fig,axesplt.subplots(1,len(models),figsize(16,5))x_min,x_maxX2[:,0].min()-0.5,X2[:,0].max()0.5y_min,y_maxX2[:,1].min()-0.5,X2[:,1].max()0.5xx,yynp.meshgrid(np.linspace(x_min,x_max,300),np.linspace(y_min,y_max,300))forax,(name,_)inzip(axes,models.items()):clfmodels[name]clf.fit(X_tr2,y_tr2)Zclf.predict(np.c_[xx.ravel(),yy.ravel()]).reshape(xx.shape)ax.contourf(xx,yy,Z,alpha0.3,cmapplt.cm.Set1)scatterax.scatter(X_te2[:,0],X_te2[:,1],cy_te2,cmapplt.cm.Set1,edgecolorsk,s40)acc2clf.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(鸢尾花分类决策边界对比仅用前两个特征,fontsize14)plt.tight_layout()fig_pathos.path.join(BASE_DIR,iris_decision_boundary.png)plt.savefig(fig_path,dpi120)print(f\n决策边界对比图已保存到:{fig_path})# ---- 用全特征模型做一个示例预测 ----best_namemax(results,keyresults.get)best_modelmodels[best_name]samplenp.array([[5.1,3.5,1.4,0.2]])# 典型山鸢尾predbest_model.predict(sample)[0]print(f\n示例预测{best_name}: 样本{sample.tolist()}-{target_names[pred]})if__name____main__:main()
返回列表