决策树分类实战:从Iris数据集入门到模型优化

1次阅读
没有评论

共计 3184 个字符,预计需要花费 8 分钟才能阅读完成。

image.webp

背景与数据集介绍

Iris 数据集是机器学习领域最经典的入门数据集之一,由统计学家 Fisher 在 1936 年首次使用。这个数据集包含了 150 个鸢尾花样本,每个样本有 4 个特征:花萼长度 (sepal length)、花萼宽度(sepal width)、花瓣长度(petal length) 和花瓣宽度 (petal width)。这些样本分为 3 类,每类 50 个样本:山鸢尾(Iris-setosa)、变色鸢尾(Iris-versicolor) 和维吉尼亚鸢尾(Iris-virginica)。

决策树分类实战:从 Iris 数据集入门到模型优化

这个数据集特别适合新手学习分类算法,因为:

  • 数据量小但特征明显,容易可视化理解
  • 特征都是数值型,无需复杂预处理
  • 三类样本完全平衡,不会出现类别不平衡问题
  • 特征与类别间有明显的相关性,模型容易学习

决策树原理通俗解释

决策树就像玩 20 个问题游戏:通过一系列是 / 否问题逐步缩小范围,最终得到答案。在机器学习中,决策树算法会自动学习这些问题的顺序和内容。

决策树的核心是如何选择最佳分裂点,常用两个指标:

  1. 信息增益 :衡量一个特征能带来多少 ” 信息量 ”。就像考试时,能最大程度区分学生水平的题目就是好题目。数学上基于熵(entropy) 计算,熵越小表示纯度越高。

  2. 基尼系数:类似信息增益,但计算更简单。衡量一个随机选中的样本被错误分类的概率,值越小越好。

决策树会递归地选择能使信息增益最大 (或基尼系数最小) 的特征进行分裂,直到满足停止条件(如达到最大深度或样本数太少)。

完整实现流程

1. 数据加载与预处理

首先导入必要库并加载数据:

from sklearn.datasets import load_iris
from sklearn.tree import DecisionTreeClassifier, export_graphviz
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score, confusion_matrix
import matplotlib.pyplot as plt
import seaborn as sns

# 加载数据
iris = load_iris()
X = iris.data  # 特征
Y = iris.target  # 标签
feature_names = iris.feature_names  # 特征名
class_names = iris.target_names  # 类别名

# 查看数据形状
print(f"特征数据形状: {X.shape}")
print(f"标签数据形状: {Y.shape}")

2. 数据探索与可视化

在建模前先了解数据分布:

# 绘制特征分布
plt.figure(figsize=(12, 8))
for i in range(4):
    plt.subplot(2, 2, i+1)
    sns.histplot(X[:, i], kde=True)
    plt.title(feature_names[i])
plt.tight_layout()
plt.show()

# 绘制特征间关系
sns.pairplot(sns.load_dataset("iris"), hue="species")
plt.show()

3. 划分训练集和测试集

# 随机划分训练集和测试集(7:3 比例)
X_train, X_test, Y_train, Y_test = train_test_split(X, Y, test_size=0.3, random_state=42)

4. 模型训练与评估

# 创建决策树分类器
clf = DecisionTreeClassifier(random_state=42, max_depth=3)

# 训练模型
clf.fit(X_train, Y_train)

# 预测测试集
Y_pred = clf.predict(X_test)

# 计算准确率
accuracy = accuracy_score(Y_test, Y_pred)
print(f"模型准确率: {accuracy:.2f}")

# 绘制混淆矩阵
cm = confusion_matrix(Y_test, Y_pred)
plt.figure(figsize=(6, 6))
sns.heatmap(cm, annot=True, fmt='d', 
            xticklabels=class_names, 
            yticklabels=class_names)
plt.xlabel('预测标签')
plt.ylabel('真实标签')
plt.title('混淆矩阵')
plt.show()

5. 可视化决策树

# 导出决策树图形
import graphviz
dot_data = export_graphviz(clf, out_file=None, 
                         feature_names=feature_names,  
                         class_names=class_names,  
                         filled=True, rounded=True,  
                         special_characters=True)
graph = graphviz.Source(dot_data)
graph.render("iris_tree")  # 保存为 PDF

graph  # 在 Jupyter 中显示

模型优化与调参

决策树有几个关键参数需要调整:

  1. max_depth:树的最大深度。太深容易过拟合,太浅可能欠拟合。建议从 3 开始尝试。

  2. min_samples_split:节点分裂所需的最小样本数。可以防止树过于复杂。

  3. min_samples_leaf:叶节点所需的最小样本数。较大的值有正则化效果。

  4. criterion:分裂标准,可选 ”gini” 或 ”entropy”。通常差异不大。

网格搜索找最优参数:

from sklearn.model_selection import GridSearchCV

param_grid = {'max_depth': [2, 3, 4, 5],
    'min_samples_split': [2, 5, 10],
    'min_samples_leaf': [1, 2, 4]
}

grid_search = GridSearchCV(DecisionTreeClassifier(random_state=42),
                          param_grid, cv=5)
grid_search.fit(X_train, Y_train)

print("最佳参数:", grid_search.best_params_)
print("最佳交叉验证分数: {:.2f}".format(grid_search.best_score_))

避坑指南

新手常犯的错误及解决方案:

  1. 忘记设置 random_state:决策树的训练具有随机性,不设置会导致每次结果不同。解决方案:固定 random_state 参数。

  2. 不进行特征缩放:虽然决策树不需要特征缩放,但很多其他算法需要。养成先查看数据分布的习惯。

  3. 忽略过拟合:决策树很容易过拟合。解决方案:

  4. 使用预剪枝(限制 max_depth 等参数)
  5. 使用后剪枝(ccp_alpha 参数)
  6. 查看训练集和测试集表现的差距

  7. 不看特征重要性:决策树可以提供特征重要性评分。解决方案:

    plt.barh(feature_names, clf.feature_importances_)
    plt.title("特征重要性")
    plt.show()

延伸思考

  1. 如果新增一个与类别无关的随机特征,决策树会如何处理?对模型性能有什么影响?
  2. 为什么决策树在 Iris 数据集上表现这么好?如果特征间相关性很低,决策树还适用吗?
  3. 如何将决策树扩展为随机森林?随机森林相比单棵决策树有什么优势?

结语

通过这个小项目,我们完成了从数据加载到模型优化的完整流程。决策树最大的优势是直观易懂,非常适合作为第一个学习的机器学习算法。虽然 Iris 数据集很简单,但掌握这些基础技能后,你可以用同样的方法处理更复杂的数据集。建议下一步尝试在 UCI 机器学习库中找一个新数据集,独立完成类似的分类任务。

正文完
 0
评论(没有评论)