用Python+Sklearn玩转决策树:鸢尾花分类实战(附完整代码与数据集)

鸢尾花分类是机器学习领域的经典入门项目,就像程序员界的"Hello World"。但不同于简单的打印语句,这个项目能让你真正触摸到机器学习的核心——如何让计算机从数据中学习规律。今天,我们就用Python和Sklearn库,从零开始构建一个决策树分类器,不仅会写代码,更要理解每一行代码背后的逻辑。

1. 环境准备与数据探索

在开始之前,确保你的Python环境已经安装了以下库:

pip install numpy pandas matplotlib scikit-learn

鸢尾花数据集是Sklearn自带的经典数据集,包含三种鸢尾花(山鸢尾、变色鸢尾、维吉尼亚鸢尾)的四个特征测量值:

  • 花萼长度(sepal length)
  • 花萼宽度(sepal width)
  • 花瓣长度(petal length)
  • 花瓣宽度(petal width)

让我们先加载并查看数据:

from sklearn.datasets import load_iris
import pandas as pd

iris = load_iris()
df = pd.DataFrame(iris.data, columns=iris.feature_names)
df['target'] = iris.target
df['flower_name'] = df['target'].apply(lambda x: iris.target_names[x])

print(df.head())
print("\n数据统计描述:")
print(df.describe())

你会看到类似这样的输出:

   sepal length (cm)  sepal width (cm)  ...  target    flower_name
0                5.1               3.5  ...       0  setosa
1                4.9               3.0  ...       0  setosa
2                4.7               3.2  ...       0  setosa
3                4.6               3.1  ...       0  setosa
4                5.0               3.6  ...       0  setosa

数据统计描述:
       sepal length (cm)  sepal width (cm)  ...  petal width (cm)      target
count         150.000000        150.000000  ...        150.000000  150.000000
mean            5.843333          3.057333  ...          1.199333    1.000000
std             0.828066          0.435866  ...          0.762238    0.819232
min             4.300000          2.000000  ...          0.100000    0.000000
25%             5.100000          2.800000  ...          0.300000    0.000000
50%             5.800000          3.000000  ...          1.300000    1.000000
75%             6.400000          3.300000  ...          1.800000    2.000000
max             7.900000          4.400000  ...          2.500000    2.000000

提示:数据探索是机器学习项目的重要第一步,通过describe()方法可以快速了解数据的分布情况,这对后续的特征工程和模型选择都有指导意义。

2. 决策树模型构建

决策树的核心思想是通过一系列的判断规则对数据进行分类。在Sklearn中,DecisionTreeClassifier类提供了完整的决策树实现:

from sklearn.tree import DecisionTreeClassifier
from sklearn.model_selection import train_test_split

# 准备数据
X = df[iris.feature_names]
y = df['target']

# 划分训练集和测试集
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)

# 创建决策树模型
clf = DecisionTreeClassifier(criterion='gini', max_depth=3, 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:.2%}")
print(f"测试集准确率: {test_score:.2%}")

几个关键参数解释:

参数 说明 常用值
criterion 分裂质量的衡量标准 'gini'(基尼系数)或'entropy'(信息增益)
max_depth 树的最大深度 整数,None表示不限制
min_samples_split 节点分裂所需最小样本数 默认2
min_samples_leaf 叶节点所需最小样本数 默认1
random_state 随机种子 任意整数,保证结果可复现

3. 决策树可视化与解释

理解决策树的工作方式比单纯使用它更重要。Sklearn提供了可视化工具:

from sklearn.tree import plot_tree
import matplotlib.pyplot as plt

plt.figure(figsize=(20,10))
plot_tree(clf, 
          feature_names=iris.feature_names,
          class_names=iris.target_names,
          filled=True, 
          rounded=True)
plt.show()

这张图展示了决策树如何做出判断:

  1. 首先检查花瓣宽度(petal width)是否≤0.8cm
    • 如果是,则分类为山鸢尾(setosa)
    • 如果不是,继续检查花瓣宽度是否≤1.75cm
      • 如果是,继续检查花瓣长度(petal length)是否≤4.95cm
        • 如果是,分类为变色鸢尾(versicolor)
        • 如果不是,分类为维吉尼亚鸢尾(virginica)
      • 如果不是,分类为维吉尼亚鸢尾(virginica)

注意:可视化时filled=True参数会根据类别着色,颜色深浅表示纯度,越深表示该节点中某一类占比越高。

4. 模型调优与验证

默认参数可能不是最优的,我们需要通过交叉验证寻找最佳参数:

from sklearn.model_selection import GridSearchCV

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

grid_search = GridSearchCV(DecisionTreeClassifier(random_state=42),
                          param_grid,
                          cv=5,
                          scoring='accuracy')
grid_search.fit(X_train, y_train)

print("最佳参数:", grid_search.best_params_)
print("最佳得分:", grid_search.best_score_)

best_clf = grid_search.best_estimator_
print("测试集准确率:", best_clf.score(X_test, y_test))

调优后,我们还可以通过特征重要性了解哪些特征对分类贡献最大:

importances = best_clf.feature_importances_
feature_importance = pd.DataFrame({
    'feature': iris.feature_names,
    'importance': importances
}).sort_values('importance', ascending=False)

print(feature_importance)

典型输出可能如下:

               feature  importance
2   petal length (cm)    0.564860
3    petal width (cm)    0.425140
0  sepal length (cm)    0.010000
1   sepal width (cm)    0.000000

这表明花瓣尺寸比花萼尺寸对分类更重要,特别是花瓣长度。

5. 决策树的优缺点与实际应用

决策树在实际项目中有其独特的优势和局限性:

优势:

  • 直观易懂,决策过程可视化
  • 对数据预处理要求低,能处理数值和类别特征
  • 不需要特征缩放
  • 可以处理多输出问题

局限性:

  • 容易过拟合,需要剪枝
  • 对数据的小变化敏感
  • 倾向于选择特征数量多的分割
  • 不适合处理线性可分数据

在实际应用中,决策树常用于:

  • 客户分群与画像
  • 信用风险评估
  • 医疗诊断辅助
  • 工业生产质量控制

6. 进阶技巧与扩展

掌握了基础后,可以尝试以下进阶技巧:

  1. 处理过拟合
# 通过预剪枝限制树生长
pruned_clf = DecisionTreeClassifier(
    max_depth=3,
    min_samples_leaf=5,
    ccp_alpha=0.02  # 代价复杂度剪枝参数
)
  1. 处理类别不平衡
# 使用类别权重
balanced_clf = DecisionTreeClassifier(
    class_weight='balanced'
)
  1. 导出决策规则
from sklearn.tree import export_text

tree_rules = export_text(best_clf, 
                        feature_names=list(iris.feature_names))
print(tree_rules)
  1. 集成学习方法
from sklearn.ensemble import RandomForestClassifier

rf_clf = RandomForestClassifier(
    n_estimators=100,
    max_depth=3,
    random_state=42
)
rf_clf.fit(X_train, y_train)

决策树虽然简单,但它是许多强大算法(如随机森林、梯度提升树)的基础。理解它如何工作,将为学习更复杂的模型打下坚实基础。

Logo

这里是“一人公司”的成长家园。我们提供从产品曝光、技术变现到法律财税的全栈内容,并连接云服务、办公空间等稀缺资源,助你专注创造,无忧运营。

更多推荐