共计 3184 个字符,预计需要花费 8 分钟才能阅读完成。
背景与数据集介绍
Iris 数据集是机器学习领域最经典的入门数据集之一,由统计学家 Fisher 在 1936 年首次使用。这个数据集包含了 150 个鸢尾花样本,每个样本有 4 个特征:花萼长度 (sepal length)、花萼宽度(sepal width)、花瓣长度(petal length) 和花瓣宽度 (petal width)。这些样本分为 3 类,每类 50 个样本:山鸢尾(Iris-setosa)、变色鸢尾(Iris-versicolor) 和维吉尼亚鸢尾(Iris-virginica)。

这个数据集特别适合新手学习分类算法,因为:
- 数据量小但特征明显,容易可视化理解
- 特征都是数值型,无需复杂预处理
- 三类样本完全平衡,不会出现类别不平衡问题
- 特征与类别间有明显的相关性,模型容易学习
决策树原理通俗解释
决策树就像玩 20 个问题游戏:通过一系列是 / 否问题逐步缩小范围,最终得到答案。在机器学习中,决策树算法会自动学习这些问题的顺序和内容。
决策树的核心是如何选择最佳分裂点,常用两个指标:
-
信息增益 :衡量一个特征能带来多少 ” 信息量 ”。就像考试时,能最大程度区分学生水平的题目就是好题目。数学上基于熵(entropy) 计算,熵越小表示纯度越高。
-
基尼系数:类似信息增益,但计算更简单。衡量一个随机选中的样本被错误分类的概率,值越小越好。
决策树会递归地选择能使信息增益最大 (或基尼系数最小) 的特征进行分裂,直到满足停止条件(如达到最大深度或样本数太少)。
完整实现流程
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 中显示
模型优化与调参
决策树有几个关键参数需要调整:
-
max_depth:树的最大深度。太深容易过拟合,太浅可能欠拟合。建议从 3 开始尝试。
-
min_samples_split:节点分裂所需的最小样本数。可以防止树过于复杂。
-
min_samples_leaf:叶节点所需的最小样本数。较大的值有正则化效果。
-
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_))
避坑指南
新手常犯的错误及解决方案:
-
忘记设置 random_state:决策树的训练具有随机性,不设置会导致每次结果不同。解决方案:固定 random_state 参数。
-
不进行特征缩放:虽然决策树不需要特征缩放,但很多其他算法需要。养成先查看数据分布的习惯。
-
忽略过拟合:决策树很容易过拟合。解决方案:
- 使用预剪枝(限制 max_depth 等参数)
- 使用后剪枝(ccp_alpha 参数)
-
查看训练集和测试集表现的差距
-
不看特征重要性:决策树可以提供特征重要性评分。解决方案:
plt.barh(feature_names, clf.feature_importances_) plt.title("特征重要性") plt.show()
延伸思考
- 如果新增一个与类别无关的随机特征,决策树会如何处理?对模型性能有什么影响?
- 为什么决策树在 Iris 数据集上表现这么好?如果特征间相关性很低,决策树还适用吗?
- 如何将决策树扩展为随机森林?随机森林相比单棵决策树有什么优势?
结语
通过这个小项目,我们完成了从数据加载到模型优化的完整流程。决策树最大的优势是直观易懂,非常适合作为第一个学习的机器学习算法。虽然 Iris 数据集很简单,但掌握这些基础技能后,你可以用同样的方法处理更复杂的数据集。建议下一步尝试在 UCI 机器学习库中找一个新数据集,独立完成类似的分类任务。
