决策树
分类算法 - 决策树
Section titled “分类算法 - 决策树”决策树(Decision Trees)简介
Section titled “决策树(Decision Trees)简介”决策树是一种多用途的监督学习(supervised learning)算法,可用于分类(classification)和回归(regression)任务。它通过基于特征值递归地划分数据来工作,创建一个树状结构,其中:
- 内部节点(Internal Nodes): 代表对特定特征的测试(例如,‘血糖水平 > 120吗?’)。
- 分支(Branches): 代表测试的结果(例如,‘是’ 或 ‘否’)。
- 叶节点(Leaf Nodes,或终端节点 Terminal Nodes): 代表最终的预测结果(分类中的类别标签,或回归中的连续值)。
它们之所以受欢迎,是因为相对容易理解和解释,模仿了人类的决策过程。
从概念上讲,想象一个用于对水果进行分类的流程图。一个节点可能会问 ‘颜色是红色的吗?’ 如果是,一个分支会导向另一个节点 ‘大小 < 5cm吗?’ 如果是,一个叶节点可能会预测 ‘樱桃’。如果不是,另一个叶节点可能会预测 ‘苹果’。
构建决策树:关键概念
Section titled “构建决策树:关键概念”核心挑战在于确定每个节点的最佳特征和分裂点(split point),以有效地划分数据。这通常通过最大化**信息增益(information gain)或最小化不纯度(impurity)**来完成。
不纯度度量(Impurity Measures)
Section titled “不纯度度量(Impurity Measures)”不纯度度量量化了节点中类别标签的“混合程度”。一个只包含一个类别的节点不纯度为零(它是纯的)。
- 基尼不纯度(Gini Impurity): 度量在节点中随机选择一个元素并根据节点中标签的分布随机为其标记类别时,被错误分类的概率。范围从 0(纯)到 0.5(对于 2 个类别,不纯度最大)。公式:
Gini = 1 - Σ(pᵢ)²,其中pᵢ是节点处类别i的实例比例。 - 熵(Entropy): 度量节点中的不确定性或随机性。范围从 0(纯)到
log₂(num_classes)。公式:Entropy = - Σ(pᵢ * log₂(pᵢ))。
像 CART(Classification and Regression Trees)这样的常见算法通常使用基尼不纯度,而 ID3 和 C4.5 使用熵/信息增益。
分裂标准(Splitting Criteria)
Section titled “分裂标准(Splitting Criteria)”在每个节点,算法会考虑所有特征的潜在分裂点。对于给定的特征,它可能会测试各种分裂点(对于连续特征)或类别(对于分类特征)。
最好的分裂是导致不纯度减少(reduction in impurity)最大(或信息增益最高,基于熵减少)的分裂。增益计算为父节点的不纯度减去分裂产生的子节点的加权平均不纯度。
停止标准与剪枝(Stopping Criteria & Pruning)
Section titled “停止标准与剪枝(Stopping Criteria & Pruning)”如果允许无限生长,决策树可以完美地拟合训练数据,导致过拟合(overfitting)。为了防止这种情况,我们需要停止标准或剪枝:
- 最大深度(Maximum Depth,
max_depth): 限制树的最大层数。 - 每个分裂所需的最小样本数(Minimum Samples per Split,
min_samples_split): 要求节点在分裂前至少包含指定数量的数据点。 - 每个叶节点所需的最小样本数(Minimum Samples per Leaf,
min_samples_leaf): 要求每个叶节点至少包含指定数量的数据点。 - 最大叶节点数(Maximum Leaf Nodes,
max_leaf_nodes): 限制叶节点的总数。 - 最小不纯度减少量(Minimum Impurity Decrease,
min_impurity_decrease): 仅当不纯度减少量超过阈值时才进行分裂。 - 剪枝(Pruning): 首先构建完整的树,然后移除预测能力很小的分支(例如,成本复杂度剪枝 Cost Complexity Pruning)。
这些是需要调优的超参数(hyperparameters)(例如,使用 GridSearchCV)。
在 Python 中的实现(Scikit-learn)
Section titled “在 Python 中的实现(Scikit-learn)”让我们使用 Scikit-learn 在 Pima Indians Diabetes 数据集上实现一个 DecisionTreeClassifier。
import pandas as pdfrom sklearn.tree import DecisionTreeClassifier, plot_tree # Updated import for visualizationfrom sklearn.model_selection import train_test_splitfrom sklearn.metrics import classification_report, confusion_matrix, accuracy_scorefrom sklearn.datasets import fetch_openmlimport matplotlib.pyplot as plt
# 1. Load Pima Indians Diabetes datasetpima = fetch_openml(name='diabetes', version=1, as_frame=True, parser='pandas')df = pima.framedf.columns = ['num_pregnancies', 'glucose', 'bp', 'skin', 'insulin', 'bmi', 'pedigree', 'age', 'class']# Ensure class is numericdf['class'] = pd.to_numeric(df['class'])
# Define features and targetfeature_cols = ['num_pregnancies', 'insulin', 'bmi', 'age', 'glucose', 'bp', 'pedigree']X = df[feature_cols]y = df['class']
# 2. Split dataset into training set and test set# 70% training and 30% testX_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3, random_state=42)
# 3. Create Decision Tree Classifier object# We can set hyperparameters here, e.g., max_depth=5 to prevent overfittingclf = DecisionTreeClassifier(max_depth=5, random_state=42)
# 4. Train Decision Tree Classifierclf = clf.fit(X_train, y_train)
# 5. Predict the response for test datasety_pred = clf.predict(X_test)
# 6. Evaluate Modelprint("--- Model Evaluation ---")conf_mat = confusion_matrix(y_test, y_pred)print("Confusion Matrix:")print(conf_mat)
class_rep = classification_report(y_test, y_pred)print("\nClassification Report:")print(class_rep)
acc_score = accuracy_score(y_test, y_pred)print(f"\nAccuracy: {acc_score:.4f}")
# 7. Visualize the Decision Tree (optional, requires matplotlib)plt.figure(figsize=(20,10))plot_tree(clf, filled=True, rounded=True, class_names=['No Diabetes (0)', 'Diabetes (1)'], # Use string names for clarity feature_names=feature_cols, fontsize=10)plt.title("Decision Tree Visualization (max_depth=5)")# plt.savefig('pima_diabetes_tree.png') # Optional: save the figureplt.show()输出(示例)
Section titled “输出(示例)”--- Model Evaluation ---Confusion Matrix:[[122 29] [ 37 43]]
Classification Report: precision recall f1-score support
0 0.77 0.81 0.79 151 1 0.60 0.54 0.57 80
accuracy 0.71 231 macro avg 0.68 0.67 0.68 231weighted avg 0.71 0.71 0.71 231
Accuracy: 0.7143输出显示了最大深度为 5 的决策树的性能指标。准确率(Accuracy)约为 71.4%。混淆矩阵(Confusion Matrix)和分类报告(Classification Report)提供了按类别划分的更详细性能信息。代码还生成了决策树结构的图。(该图将显示一个树状图,节点显示分裂标准、不纯度、样本数和预测类别)。
特征重要性(Feature Importance)
Section titled “特征重要性(Feature Importance)”决策树(以及基于树的集成方法)可以提供特征重要性(Feature Importance)的估计,指示哪些特征在树中的所有分裂中对数据划分最有帮助。这通常基于该特征贡献的总不纯度减少量。
# --- Feature Importance ---importances = clf.feature_importances_feature_importance_df = pd.DataFrame({'feature': feature_cols, 'importance': importances})feature_importance_df = feature_importance_df.sort_values('importance', ascending=False)
print("\n--- Feature Importances ---")print(feature_importance_df)输出(示例)
Section titled “输出(示例)”--- Feature Importances --- feature importance4 glucose 0.3852 bmi 0.1883 age 0.1406 pedigree 0.1110 num_pregnancies 0.0775 bp 0.0651 insulin 0.034在此示例运行中,根据这棵特定的树,‘glucose’(血糖)、‘bmi’(身体质量指数)和 ‘age’(年龄)似乎是最重要的特征。
- 可解释性(Interpretability): 易于理解和可视化决策过程。
- 最少数据准备(Minimal Data Prep): 通常需要较少的数据准备(例如,不需要特征缩放 feature scaling)。可以处理数值和分类数据(尽管 Scikit-learn 需要数值输入)。
- 非线性关系(Non-linear Relationships): 可以捕捉数据中复杂的非线性模式。
- 特征重要性(Feature Importance): 内在地执行特征选择并提供重要性得分。
- 过拟合(Overfitting): 容易创建过于复杂的树,泛化能力差。需要仔细调优(例如,
max_depth)或剪枝。 - 不稳定性(Instability): 数据中的微小变化可能导致完全不同的树结构。
- 偏差(Bias): 如果某些类别占主导地位,可能会创建有偏差的树。
- 最优性(Optimality): 找到全局最优决策树在计算上是不可行的(NP-hard),因此算法使用贪婪方法(局部最优分裂),这可能无法产生整体最佳的树。
由于过拟合和不稳定性的问题,单一决策树在实践中常常被集成方法(ensemble methods)取代,如随机森林(Random Forests)或梯度提升树(Gradient Boosting Trees),这些方法基于决策树的概念构建。
- Scikit-learn 决策树文档: https://scikit-learn.org/stable/modules/tree.html
- StatQuest: 决策树解释: https://statquest.org/video-index/ (搜索 Decision Trees)