入门实践工程一:基于 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)

安装依赖

bash 复制代码
pip install scikit-learn matplotlib numpy

注意事项:

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

验证安装

python 复制代码
import sklearn
print(f"scikit-learn 版本: {sklearn.__version__}")
# 应该输出类似: scikit-learn 版本: 1.3.0

运行方式

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

bash 复制代码
python main.py

方法二:使用 requirements.txt

bash 复制代码
pip install -r requirements.txt
python main.py

运行过程解析

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

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

代码详解

1. 数据加载与探索

python 复制代码
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. 数据划分策略

python 复制代码
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. 模型定义与训练

python 复制代码
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. 决策边界可视化原理

python 复制代码
# 创建网格点
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 失败

解决方案:

bash 复制代码
# 使用国内镜像源
pip install scikit-learn matplotlib numpy -i https://pypi.tuna.tsinghua.edu.cn/simple

# 或使用 conda
conda install scikit-learn matplotlib numpy

Q2: 运行时报错 "ModuleNotFoundError"

可能原因: 依赖包未正确安装

解决方案:

bash 复制代码
# 检查已安装的包
pip list | grep -E "scikit-learn|matplotlib|numpy"

# 重新安装
pip uninstall scikit-learn matplotlib numpy
pip install scikit-learn matplotlib numpy

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

解决方案:

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

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

原因: 未设置随机种子

解决方案: 代码中已设置 random_state=42,确保结果可复现

扩展方向与进阶学习

1. 特征工程扩展

python 复制代码
# 添加特征标准化
from sklearn.preprocessing import StandardScaler
scaler = StandardScaler()
X_scaled = scaler.fit_transform(X)

# 添加 PCA 降维可视化
from sklearn.decomposition import PCA
pca = PCA(n_components=2)
X_pca = pca.fit_transform(X)

2. 更换数据集挑战

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

3. 模型扩展与对比

python 复制代码
# 添加逻辑回归
from sklearn.linear_model import LogisticRegression
models["逻辑回归"] = LogisticRegression(max_iter=1000)

# 添加 XGBoost
from xgboost import XGBClassifier
models["XGBoost"] = XGBClassifier(n_estimators=100)

# 绘制 ROC 曲线(二分类)
from sklearn.metrics import roc_curve, auc
fpr, tpr, _ = roc_curve(y_test_binary, y_score)
roc_auc = auc(fpr, tpr)

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

python 复制代码
from sklearn.model_selection import cross_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. 模型可解释性

python 复制代码
# 随机森林特征重要性
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] for i in indices], 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

python 复制代码
"""
入门实践工程一:基于 sklearn 的鸢尾花分类(传统机器学习入门)
=========================================================
使用 scikit-learn 内置的鸢尾花数据集,训练并对比 KNN / SVM / 随机森林
三种经典分类器,评估准确率,并绘制二维决策边界可视化图。
全程无需深度学习框架、无需联网,是最轻量的 AI 入门项目。

运行:
    python main.py
"""

import os

import matplotlib

matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
from sklearn.datasets import load_iris
from sklearn.ensemble import RandomForestClassifier
from sklearn.model_selection import train_test_split
from sklearn.neighbors import KNeighborsClassifier
from sklearn.svm import SVC

BASE_DIR = os.path.dirname(os.path.abspath(__file__))


def main():
    iris = load_iris()
    X, y = iris.data, iris.target
    feature_names = iris.feature_names
    target_names = iris.target_names
    print(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 = {}
    for name, clf in models.items():
        clf.fit(X_train, y_train)
        acc = clf.score(X_test, y_test)
        results[name] = acc
        print(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.5
    y_min, y_max = X2[:, 1].min() - 0.5, X2[:, 1].max() + 0.5
    xx, yy = np.meshgrid(np.linspace(x_min, x_max, 300),
                         np.linspace(y_min, y_max, 300))

    for ax, (name, _) in zip(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()
相关推荐
TechEdu2026061 小时前
[人工智能]Scikit-learn:Python中的实用机器学习库
人工智能·机器学习·ai·scikit-learn
databook2 小时前
基于模型的重要性评分进行“特征排序”
python·机器学习·scikit-learn
404NotFOund2 小时前
小白本地部署微调耍起
机器学习·开源
具身新纪元4 小时前
CVPR 2026|新缺陷不断上线,如何让工业视觉检测模型不忘旧账?
人工智能·深度学习·目标检测·机器学习·计算机视觉·目标跟踪·视觉检测
only-qi4 小时前
大模型微调流程深度解析:从面试题到工程实践
人工智能·机器学习·ai·llm
一次旅行5 小时前
Ollama本地私有化大模型完整工程实战
人工智能·机器学习·github
小趴蔡ha13 小时前
02 Anaconda、Python 与机器学习开发环境入门
python·机器学习·anaconda
小趴蔡ha19 小时前
04 Pandas 数据清洗实战:从表格读取到机器学习特征准备
人工智能·机器学习·pandas
SomeB1oody19 小时前
【RustyML入门】2.6. 线性判别分析
开发语言·后端·机器学习·rust·教程