1. 项目概述:从数据到洞察的经典旅程

“鸢尾花分类”这个项目,对于任何一个刚踏入机器学习或数据分析领域的朋友来说,都像是一本必读的入门经典。它简单、清晰,却又完整地串联起了从数据理解、探索分析、模型构建到结果可视化的全流程。这个项目标题——“鸢尾花分类与直方图、散点图的绘制及可视化决策树”——精准地概括了数据分析与机器学习入门阶段的核心技能闭环。

简单来说,我们手头有一份关于鸢尾花的数据集,里面记录了三种不同品种的鸢尾花(山鸢尾、变色鸢尾、维吉尼亚鸢尾)的四个关键特征:萼片长度、萼片宽度、花瓣长度、花瓣宽度。我们的目标,就是教会计算机如何根据这四个测量值,自动判断一朵未知的鸢尾花属于哪个品种。这听起来像是一个分类任务,没错,它正是监督学习中最基础的分类问题。

但这个过程远不止是“丢给模型,得出结果”那么简单。标题中提到的“直方图、散点图的绘制”是整个流程中至关重要的一环,我们称之为 探索性数据分析 。在盲目建模之前,我们必须先“认识”我们的数据:每个特征的分布情况如何?不同特征之间有什么关系?不同类别的花在这些特征上是否有明显的区分?直方图能告诉我们单个特征的数值分布,而散点图则能揭示两个特征之间的关联与类别分离情况。这些图表是我们与数据对话的语言,能帮助我们形成初步的直觉,甚至发现数据中的异常。

最后,“可视化决策树”则是将模型“黑箱”透明化的关键一步。决策树是一种非常直观的机器学习算法,它的决策逻辑就像是一系列“如果…那么…”的判断规则。将其可视化出来,我们不仅能评估模型的分类效果,更能清晰地理解模型是依据哪些特征、在什么阈值上做出了分类决策。这对于模型的可解释性至关重要,尤其是在向业务方解释模型逻辑时,一张清晰的决策树图胜过千言万语。

因此,这个项目非常适合以下几类朋友:刚接触Python数据分析库(如Pandas, Matplotlib, Seaborn)的新手,希望系统学习机器学习工作流程的初学者,以及需要向他人清晰展示分析过程和模型逻辑的从业者。接下来,我将以一个从业者的视角,带你完整走一遍这个经典项目,并分享一些实操中容易踩坑的细节和技巧。

2. 环境准备与数据初探

工欲善其事,必先利其器。我们首先需要搭建好分析环境。最经典和便捷的组合是使用Python的Jupyter Notebook或JupyterLab,配合几个核心的数据科学库。

2.1 核心工具库安装与导入

通常,我们会使用Anaconda来管理Python环境,它能很好地解决库之间的依赖问题。如果你使用的是纯净的Python,可以通过pip安装以下核心库:

pip install numpy pandas matplotlib seaborn scikit-learn

安装完成后,在Notebook或脚本的开头,我们惯例性地导入这些库,并通常给它们起一个简短的别名,方便后续调用。

# 基础数据处理与计算
import numpy as np
import pandas as pd

# 数据可视化
import matplotlib.pyplot as plt
import seaborn as sns

# 机器学习
from sklearn import datasets
from sklearn.model_selection import train_test_split
from sklearn.tree import DecisionTreeClassifier, plot_tree
from sklearn.metrics import classification_report, confusion_matrix

# 设置可视化样式(可选,让图表更好看)
sns.set_style("whitegrid")
plt.rcParams['font.sans-serif'] = ['SimHei']  # 用来正常显示中文标签
plt.rcParams['axes.unicode_minus'] = False  # 用来正常显示负号

注意 :中文字体设置( SimHei )在Windows系统下通常有效。如果你在Mac或Linux上运行,可能需要替换为其他已安装的中文字体名,如 ‘Arial Unicode MS’ ,或者移除这两行代码使用默认英文字体,避免出现乱码方块。

2.2 加载与审视鸢尾花数据集

Scikit-learn库内置了许多经典数据集,鸢尾花(Iris)就是其中之一,这为我们省去了数据收集和清洗的麻烦,可以专注于分析和建模流程。

# 加载数据集
iris = datasets.load_iris()

# 通常我们会将数据转换为Pandas的DataFrame,这样查看和处理起来更直观
iris_df = pd.DataFrame(data=iris.data, columns=iris.feature_names)
# 添加目标列(花的品种)
iris_df['target'] = iris.target
# 为了更直观,我们可以将数字标签映射回花的品种名称
iris_df['species'] = iris_df['target'].apply(lambda x: iris.target_names[x])

# 查看数据前5行
print(iris_df.head())

运行上述代码,你会看到一个结构清晰的表格。每一行代表一朵花的测量记录,前四列是特征(单位是厘米), target 是数字标签(0, 1, 2), species 是对应的品种名。

在深入分析前,我们必须先对数据有一个整体的把握:

# 查看数据集的基本信息
print(f"数据集形状: {iris_df.shape}")  # 输出 (150, 6),表示150个样本,6列(4特征+1数字标签+1品种名)
print("\n数据类型和非空值检查:")
print(iris_df.info())
print("\n数据摘要统计(针对数值列):")
print(iris_df.describe())
print("\n品种分布:")
print(iris_df['species'].value_counts())

df.info() 能告诉我们是否有缺失值(幸运的是,Iris数据集是完整的),以及每列的数据类型。 df.describe() 则提供了数值特征(萼片长宽、花瓣长宽)的统计摘要,包括均值、标准差、最小值、四分位数和最大值。这能让我们快速发现是否存在异常值(比如某个测量值远超正常范围)。

从品种分布看,三个类别各50个样本,这是一个非常均衡的数据集,这简化了我们后续的工作,不需要进行复杂的样本平衡处理。

3. 探索性数据分析:让数据自己说话

在建模之前,我们必须花时间“看”数据。探索性数据分析的目标是发现数据的内在结构、识别模式、检测异常,并初步验证我们后续可能使用的建模方法(比如线性可分性)是否适用。直方图和散点图是我们最得力的两大工具。

3.1 单变量分析:直方图与核密度估计

直方图是展示单个特征分布最直观的方式。它将数据范围划分成若干个连续的区间(称为“箱子”或“bin”),然后统计每个区间内数据点的数量。通过直方图,我们可以了解特征的集中趋势(大部分值落在哪里)、离散程度(数据是紧密还是分散)以及分布形状(是否对称,是否有偏斜)。

我们可以为每个特征单独绘制直方图,但更高效的方式是使用子图将它们放在一起对比。

# 准备特征名称列表
features = iris.feature_names
# 设置画布,1行4列,共享y轴以便比较
fig, axes = plt.subplots(1, 4, figsize=(16, 4), sharey=True)

# 循环绘制每个特征的直方图
for idx, feature in enumerate(features):
    ax = axes[idx]
    # 绘制直方图,bins参数控制区间的数量,edgecolor给柱子加边框,alpha设置透明度
    ax.hist(iris_df[feature], bins=15, edgecolor='black', alpha=0.7)
    ax.set_title(f'{feature}分布')
    ax.set_xlabel('测量值 (cm)')
    if idx == 0:  # 只在第一个子图设置y轴标签
        ax.set_ylabel('频数')

plt.tight_layout()  # 自动调整子图参数,使之填充整个图像区域,避免标签重叠
plt.show()

观察这四个直方图,你能立刻获得一些信息:花瓣长度和花瓣宽度的分布似乎不是单峰的,可能暗示了不同品种在这些特征上有明显差异。而萼片长度和宽度的分布则相对集中和对称。

为了更平滑地展示分布,并同时观察不同品种的分布差异,我们可以使用 核密度估计图 ,并按照品种进行区分。Seaborn库的 displot kdeplot 函数可以非常优雅地实现这一点。

plt.figure(figsize=(12, 8))
for i, feature in enumerate(features):
    plt.subplot(2, 2, i+1)  # 创建2行2列的子图
    # 为每个品种绘制核密度估计曲线
    for species_name in iris.target_names:
        subset = iris_df[iris_df['species'] == species_name]
        sns.kdeplot(data=subset[feature], label=species_name, fill=True, alpha=0.5)
    plt.title(f'{feature}的核密度估计')
    plt.xlabel('测量值 (cm)')
    plt.ylabel('密度')
    plt.legend()
plt.tight_layout()
plt.show()

这张图的信息量巨大。你可以清晰地看到:

  • 花瓣长度 花瓣宽度 :三个品种的分布几乎完全分开,尤其是 setosa (山鸢尾)与其他两种。这意味着仅凭这两个特征之一,就可能很好地进行分类。
  • 萼片长度 versicolor (变色鸢尾)和 virginica (维吉尼亚鸢尾)的分布有较大重叠,单独使用区分力较弱。
  • 萼片宽度 :三个品种的分布重叠较多,可能是区分能力最弱的特征。

这个分析直接影响了我们后续的特征工程和模型选择。例如,我们可能会更关注花瓣相关的特征。

3.2 双变量分析:散点图矩阵

单变量分析看个体,双变量分析看关系。散点图是研究两个连续变量之间关系的标准工具。对于多维数据,散点图矩阵可以一次性展示所有特征两两之间的关系。

最强大的方式是使用Seaborn的 pairplot 函数,它能自动绘制数据集中每对数值特征之间的散点图(对角线用直方图或KDE图替代),并且可以按类别着色。

# 使用pairplot, hue参数指定按品种着色
sns.pairplot(iris_df, hue='species', diag_kind='kde', palette='husl', height=2.5)
plt.suptitle('鸢尾花特征散点图矩阵(按品种着色)', y=1.02)  # y参数调整标题位置
plt.show()

这张图是探索性数据分析的精华所在,值得你花时间仔细研究每一个小图:

  1. 对角线(KDE图) :与之前的单变量KDE图一致,再次验证了花瓣特征的区分度。
  2. 非对角线(散点图)
    • 观察 petal length vs petal width 这个散点图:你会发现三个品种的样本点形成了三个几乎线性可分的簇。 setosa 集中在左下角(花瓣小且短), versicolor 在中间, virginica 在右上角(花瓣大且宽)。而且这两个特征呈现强烈的正相关关系(点呈带状分布)。
    • 观察 sepal length vs sepal width :点云比较分散,类别间的界限模糊。
    • 观察 petal length vs sepal length :同样能看到较好的分离,尤其是 setosa 与其他两类。

从这些散点图中,我们可以得出几个关键结论:

  • 特征重要性 :花瓣相关特征(尤其是长度和宽度)对于分类至关重要。
  • 线性可分性 :数据在由花瓣特征构成的空间中,近似线性可分,这暗示线性模型(如逻辑回归)或基于划分的模型(如决策树)可能会表现良好。
  • 特征冗余 petal length petal width 高度相关,这意味着它们提供的信息可能有所重叠。在有些模型中,我们可能需要考虑这一点,但决策树本身对特征相关性不敏感。

实操心得 pairplot 是快速了解数据集全貌的“神器”。在实际项目中,当你的特征数量不多于10个时,都可以先用它来扫一眼。如果特征太多,图会变得非常密集,这时可以考虑先进行特征选择或降维,或者只绘制与目标变量相关性最高的几个特征之间的散点图。

4. 决策树模型构建与训练

经过充分的EDA,我们对数据已经了如指掌。现在,可以开始构建分类模型了。我们选择决策树,正是因为它的可解释性与我们项目“可视化”的目标完美契合。

4.1 数据准备:划分训练集与测试集

在机器学习中,一个基本原则是: 绝不能使用测试数据来训练模型 。我们必须将数据集划分为两部分:一部分用于训练模型(训练集),另一部分用于评估模型的泛化能力(测试集)。Scikit-learn提供了方便的 train_test_split 函数。

# 准备特征矩阵X和目标向量y
X = iris_df[features]  # 只选取四个特征列
y = iris_df['target']  # 使用数字标签作为目标

# 划分数据集,test_size=0.3表示30%的数据作为测试集,random_state设置随机种子确保结果可复现
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3, random_state=42)

print(f"训练集样本数: {X_train.shape[0]}")
print(f"测试集样本数: {X_test.shape[0]}")

这里 random_state 参数非常重要。它保证了每次运行代码时,数据集的划分方式都是一样的,从而使你的实验结果可复现。在分享代码或撰写报告时,务必固定这个值。

4.2 决策树原理与关键参数

决策树通过一系列“是/否”问题对数据进行递归分割。它从根节点(全部数据)开始,选择一个特征和一个阈值,将数据分成两个子集(子节点),使得分割后的子集“纯度”更高(即同一类别的样本尽可能在一起)。这个过程在每个子节点上重复,直到满足停止条件(如达到最大深度、节点样本数过少等)。

在Scikit-learn中, DecisionTreeClassifier 有几个关键参数需要我们理解:

  • criterion :分裂质量的衡量标准。常用 gini (基尼不纯度)或 entropy (信息增益)。两者效果通常相近,Gini计算稍快。
  • max_depth :树的最大深度。这是控制模型复杂度的最重要参数。深度越大,树越复杂,越容易过拟合(在训练集上表现好,在测试集上差);深度太小,则可能欠拟合(学不到模式)。通常需要调优。
  • min_samples_split :内部节点再划分所需的最小样本数。如果某节点的样本数少于这个值,则不会继续分裂。
  • min_samples_leaf :叶节点所需的最小样本数。这个参数可以平滑模型,防止生成特别容易受噪声影响的叶节点。
  • random_state :同样用于控制随机性(如特征排序的随机性),确保可复现。

对于鸢尾花这种小数据集,我们先尝试一个简单的、不加限制的树,看看它会长成什么样。

# 初始化决策树分类器,先使用默认参数
clf = DecisionTreeClassifier(random_state=42)
# 在训练集上训练模型
clf.fit(X_train, y_train)

# 评估模型在训练集和测试集上的表现
train_score = clf.score(X_train, y_train)
test_score = clf.score(X_test, y_test)

print(f"训练集准确率: {train_score:.4f}")
print(f"测试集准确率: {test_score:.4f}")

你可能会发现,训练集准确率是100%,而测试集准确率也很高(可能在95%以上)。100%的训练准确率是一个明显的 过拟合 信号——模型完全记住了训练数据,甚至包括其中的噪声。这样的模型在未知数据上的表现可能不稳定。

4.3 模型调优:防止过拟合

我们的目标是得到一个泛化能力强的模型。因此,我们需要对决策树进行“剪枝”,控制其复杂度。最直接的方法是限制树的最大深度 max_depth

# 尝试不同的最大深度,观察模型表现
max_depths = range(1, 11)
train_scores = []
test_scores = []

for depth in max_depths:
    clf = DecisionTreeClassifier(max_depth=depth, random_state=42)
    clf.fit(X_train, y_train)
    train_scores.append(clf.score(X_train, y_train))
    test_scores.append(clf.score(X_test, y_test))

# 绘制准确率随深度变化的曲线
plt.figure(figsize=(10, 6))
plt.plot(max_depths, train_scores, 'o-', label='训练集准确率', linewidth=2)
plt.plot(max_depths, test_scores, 's-', label='测试集准确率', linewidth=2)
plt.xlabel('决策树最大深度')
plt.ylabel('准确率')
plt.title('决策树深度对模型性能的影响')
plt.legend()
plt.grid(True, linestyle='--', alpha=0.7)
plt.show()

通过这张图,你可以清晰地看到 偏差-方差权衡

  • max_depth=1 或2时(树很浅),训练集和测试集准确率都不高,这是 欠拟合 ,模型太简单,无法捕捉数据中的模式。
  • 随着深度增加,训练集准确率迅速上升至接近100%。测试集准确率也上升,但在某个点(比如深度为3或4)达到峰值。
  • 深度继续增加,训练集准确率维持在100%,但测试集准确率可能开始轻微波动或下降,这就是 过拟合 的开始。模型过于复杂,学习了训练数据中的特定噪声。

对于鸢尾花数据集, max_depth=3 4 通常是一个很好的选择,它在测试集上取得了高准确率,同时模型保持简洁。我们选择 max_depth=3 来构建最终模型。

# 使用调优后的参数重新训练模型
final_clf = DecisionTreeClassifier(max_depth=3, random_state=42)
final_clf.fit(X_train, y_train)

final_train_score = final_clf.score(X_train, y_train)
final_test_score = final_clf.score(X_test, y_test)
print(f"最终模型 - 训练集准确率: {final_train_score:.4f}")
print(f"最终模型 - 测试集准确率: {final_test_score:.4f}")

现在,我们得到了一个既准确又简洁的模型。接下来,就是揭开它“思考过程”的时刻。

5. 决策树可视化与模型解读

决策树最大的优势就是可解释性。Scikit-learn提供了 plot_tree 函数,可以方便地将训练好的树模型绘制出来。

5.1 绘制可视化决策树

# 设置图形大小
plt.figure(figsize=(20, 12))
# 绘制决策树
plot_tree(final_clf,
          feature_names=features,  # 使用特征名
          class_names=iris.target_names,  # 使用类别名
          filled=True,  # 给节点填充颜色,颜色深浅表示节点纯度
          rounded=True,  # 使用圆角矩形节点
          fontsize=12,  # 字体大小
          proportion=True)  # 显示样本比例而非具体数量
plt.title("鸢尾花分类决策树可视化 (max_depth=3)", fontsize=16)
plt.show()

这张图包含了丰富的信息,我们来学习如何解读它:

  1. 每个节点的第一行 :分裂条件。例如根节点是“ petal length (cm) <= 2.45 ”。这意味着模型首先根据花瓣长度是否小于等于2.45厘米来做决策。
  2. gini :该节点的基尼不纯度。值越接近0,表示该节点中样本的类别越纯(都属于同一类)。根节点的gini=0.667,说明三个类别混合程度较高。
  3. samples :到达该节点的总样本数。根节点是105,即我们训练集的总样本数。
  4. value :一个列表,显示该节点中每个类别的样本数量。根节点的 value = [31, 37, 37] ,对应 setosa , versicolor , virginica
  5. class :该节点中被预测的类别,即该节点中样本数最多的类别。

跟着决策路径走一遍

  • 从根节点开始:如果一朵花的花瓣长度 <= 2.45,则进入左子节点。这个节点的gini=0.0,value=[31,0,0],class=setosa。 这是一个叶节点 (没有进一步分裂)。这意味着所有花瓣长度小于等于2.45厘米的花都被直接判定为山鸢尾,而且这个判断是100%纯的(gini=0)。这与我们EDA中观察到的完全一致!
  • 如果花瓣长度 > 2.45,则进入右子节点。这个节点包含74个样本(37个versicolor, 37个virginica),gini=0.5,说明两个类别各占一半,需要进一步分裂。
  • 该节点根据 petal width (cm) <= 1.75 进行第二次分裂。如果花瓣宽度<=1.75,进入左子节点(value=[0, 34, 4]),大部分是versicolor,但混有4个virginica。
  • 这个节点继续根据 petal length (cm) <= 4.95 分裂。最终,我们得到了三个叶节点,分别对应三个类别。

5.2 特征重要性分析

决策树还可以量化每个特征在做出正确决策时的重要性。重要性是基于该特征被用于分裂节点时,所带来的不纯度减少的总量(或信息增益)来计算的。

# 获取特征重要性
importances = final_clf.feature_importances_
# 将其与特征名对应,并排序
feature_importance_df = pd.DataFrame({
    'feature': features,
    'importance': importances
}).sort_values('importance', ascending=False)

print(feature_importance_df)

# 可视化特征重要性
plt.figure(figsize=(8, 5))
sns.barplot(data=feature_importance_df, x='importance', y='feature', palette='viridis')
plt.title('决策树特征重要性排序')
plt.xlabel('重要性得分')
plt.tight_layout()
plt.show()

不出所料, petal length (cm) petal width (cm) 占据了最重要的地位,而 sepal width (cm) 的重要性为0,这意味着在这棵深度为3的树中,它根本没有被用到。这再次印证了EDA阶段我们的观察:花瓣特征提供了最强的区分信号。

注意事项 :特征重要性是 针对当前已训练的特定模型 而言的。如果一个特征没有被选入树中,其重要性为0,但这并不绝对意味着该特征与目标变量无关。它可能与其他强特征高度相关,导致其提供的信息是冗余的。在我们的案例中, sepal width 可能就是因为与 petal 特征提供的分类信息重叠,而被决策树“忽略”了。

6. 模型评估与深入洞察

准确率是一个直观的指标,但为了全面评估一个分类模型,尤其是多分类模型,我们需要更细致的工具。

6.1 混淆矩阵与分类报告

混淆矩阵是评估分类模型性能的基石。它是一个N x N的矩阵(N为类别数),行代表真实类别,列代表预测类别。对角线上的数字是预测正确的样本数,其他位置则是预测错误的样本数。

# 在测试集上进行预测
y_pred = final_clf.predict(X_test)

# 计算混淆矩阵
cm = confusion_matrix(y_test, y_pred)
print("混淆矩阵:")
print(cm)

# 使用Seaborn绘制更美观的热力图
plt.figure(figsize=(8, 6))
sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',
            xticklabels=iris.target_names,
            yticklabels=iris.target_names)
plt.title('混淆矩阵热力图')
plt.ylabel('真实标签')
plt.xlabel('预测标签')
plt.show()

从混淆矩阵中,我们可以一目了然地看到模型在哪里犯了错。例如,可能有个别 versicolor 被误判为 virginica ,或者反之。这能帮助我们定位模型的薄弱环节。

分类报告则提供了精确率、召回率、F1-score等更详细的指标。

report = classification_report(y_test, y_pred, target_names=iris.target_names)
print("分类报告:")
print(report)
  • 精确率 :在所有被预测为A类的样本中,真正是A类的比例。它关注的是 预测结果的准确性
  • 召回率 :在所有真实为A类的样本中,被模型正确预测为A类的比例。它关注的是 模型找出正样本的能力
  • F1-score :精确率和召回率的调和平均数,是一个综合指标。

对于均衡数据集,准确率通常已足够。但在类别不平衡的数据集中,精确率、召回率和F1-score更能反映模型的真实性能。

6.2 决策边界可视化(进阶)

为了更直观地理解决策树是如何在特征空间中“划界”的,我们可以绘制其决策边界。由于我们有四个特征,无法在四维空间绘图。因此,我们选取两个最重要的特征(如 petal length petal width )来绘制二维决策边界。

# 选取两个特征
X_2d = X_train[['petal length (cm)', 'petal width (cm)']].values
y_2d = y_train.values

# 训练一个仅基于这两个特征的决策树(使用相同参数)
clf_2d = DecisionTreeClassifier(max_depth=3, random_state=42)
clf_2d.fit(X_2d, y_2d)

# 创建网格来覆盖特征空间
x_min, x_max = X_2d[:, 0].min() - 0.5, X_2d[:, 0].max() + 0.5
y_min, y_max = X_2d[:, 1].min() - 0.5, X_2d[:, 1].max() + 0.5
xx, yy = np.meshgrid(np.arange(x_min, x_max, 0.02),
                     np.arange(y_min, y_max, 0.02))

# 预测网格上每个点的类别
Z = clf_2d.predict(np.c_[xx.ravel(), yy.ravel()])
Z = Z.reshape(xx.shape)

# 绘制决策区域和训练样本点
plt.figure(figsize=(12, 8))
# 绘制决策区域
plt.contourf(xx, yy, Z, alpha=0.3, cmap=plt.cm.coolwarm)
# 绘制训练数据点
scatter = plt.scatter(X_2d[:, 0], X_2d[:, 1], c=y_2d, edgecolor='black', s=50, cmap=plt.cm.coolwarm)
plt.xlabel('花瓣长度 (cm)')
plt.ylabel('花瓣宽度 (cm)')
plt.title('基于花瓣长度和宽度的决策树决策边界 (2D投影)')
plt.colorbar(scatter, ticks=[0, 1, 2], label='品种').set_ticklabels(iris.target_names)
plt.show()

这张图生动地展示了决策树如何通过一系列垂直于坐标轴的直线(因为每次分裂只基于一个特征)将特征空间划分成不同的矩形区域。每个区域对应决策树的一个叶节点,区域内的所有点都会被预测为同一个类别。你可以清晰地看到 setosa 被干净利落地分在左下角的小矩形里,而 versicolor virginica 的边界则稍微复杂一些,这正是决策树学习到的“规则”。

7. 项目总结与扩展思考

走完这个完整的流程,你已经实践了一个标准的数据分析-机器学习微型项目。让我们回顾一下核心步骤与收获:

  1. 数据加载与初探 :使用 pandas sklearn 加载数据,通过 .info() , .describe() 快速了解数据全貌,这是所有分析的基础。
  2. 探索性数据分析 :利用 matplotlib seaborn 绘制直方图、核密度图、散点图矩阵,让数据可视化,从中发现特征分布、关联和类别可分性的关键线索。 这是决定后续建模方向的关键,绝不能跳过。
  3. 数据划分 :使用 train_test_split 划分训练集和测试集,确保模型评估的公正性,并通过 random_state 保证结果可复现。
  4. 模型构建与调优 :选择决策树模型,理解其关键参数(如 max_depth ),通过绘制学习曲线进行调优,找到偏差与方差的平衡点,防止过拟合。
  5. 模型可视化与解读 :使用 plot_tree 将模型结构图形化,学习解读节点信息,理解模型的决策逻辑。同时分析特征重要性,验证EDA阶段的发现。
  6. 模型评估 :超越简单的准确率,使用混淆矩阵和分类报告进行细致评估,定位模型错误。

踩坑与心得

  • 图表的可读性 :绘制图表时,务必添加清晰的标题、轴标签和图例。使用 plt.tight_layout() 避免标签重叠。在分享或报告时,美观清晰的图表至关重要。
  • 随机种子 :在涉及随机性的步骤(如数据划分、决策树特征排序),务必设置 random_state 。这是保证实验可复现性的生命线。
  • 理解过拟合 :训练集上100%的准确率不一定是好事。一个在训练集上表现完美但在测试集上表现平平的模型,其价值远低于一个在两者上都表现良好但非完美的模型。
  • 决策树的局限性 :决策树容易过拟合,对数据的小变化敏感(高方差)。可以通过剪枝参数( max_depth , min_samples_leaf 等)控制,或者使用集成方法如随机森林来提升稳定性和性能。

扩展思考 : 这个项目虽然经典,但你可以轻松地对其进行扩展,深化学习:

  • 尝试其他模型 :用同样的数据试试逻辑回归、支持向量机或K近邻,比较它们的性能和决策边界有何不同。
  • 特征工程 :尝试创建新的特征,比如花瓣的长宽比( petal length / petal width ),看看这个新特征是否具有更强的区分能力,并重新训练决策树。
  • 超参数调优 :使用 GridSearchCV RandomizedSearchCV 对决策树的多个参数( max_depth , min_samples_split , criterion 等)进行系统性的网格搜索,寻找最优组合。
  • 处理真实数据 :寻找一个类似结构的公开数据集(如葡萄酒分类、乳腺癌诊断),但可能包含缺失值、异常值或类别不平衡。将本项目的流程应用上去,并解决这些新出现的问题。

机器学习项目就像解一道复杂的谜题,EDA是仔细观察谜面,建模是提出解题假设,评估是验证答案。鸢尾花分类项目提供了一个近乎完美的谜面,让你能专注于练习解题的完整流程。掌握了这个流程,你就拥有了应对更复杂、更真实数据挑战的基础能力。

Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐