经典决策树分类实战:从原理到Python实现

1次阅读
没有评论

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

image.webp

1. 理解决策树的基本原理

决策树是一种基于树结构的分类方法,它通过一系列的条件判断来对数据进行分类。每个内部节点代表一个特征属性上的测试,每个分支代表测试的输出,而每个叶节点代表一个类别。决策树的构建过程主要包括特征选择、决策树生成和剪枝三个步骤。

经典决策树分类实战:从原理到 Python 实现

2. 特征选择的数学基础

2.1 信息增益

信息增益是 ID3 算法中用于特征选择的准则,基于信息论中的熵概念。熵用于度量样本集合的不确定性,定义为:

$$ H(D) = -\sum_{k=1}^{K} p_k \log_2 p_k $$

其中 $p_k$ 是样本集合 D 中第 k 类样本所占的比例。对于某个特征 A,其对数据集 D 的信息增益定义为:

$$ Gain(D,A) = H(D) – \sum_{v=1}^{V} \frac{|D^v|}{|D|} H(D^v) $$

2.2 基尼系数

CART 算法使用基尼系数作为特征选择标准,它反映了从数据集 D 中随机抽取两个样本,其类别标记不一致的概率:

$$ Gini(D) = 1 – \sum_{k=1}^{K} p_k^2 $$

对于特征 A 的基尼指数定义为:

$$ Gini_index(D,A) = \sum_{v=1}^{V} \frac{|D^v|}{|D|} Gini(D^v) $$

3. 算法对比

算法 分裂标准 树类型 支持特征类型
ID3 信息增益 多叉树 离散型
C4.5 信息增益比 多叉树 离散 / 连续
CART 基尼系数 二叉树 离散 / 连续

4. Python 实现

from sklearn.tree import DecisionTreeClassifier
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
import matplotlib.pyplot as plt

def train_decision_tree(
    criterion: str = 'gini', 
    max_depth: int = None,
    random_state: int = 42
) -> DecisionTreeClassifier:
    """
    训练决策树分类器

    Parameters
    ----------
    criterion : str, optional
        分裂标准 ('gini' 或 'entropy'), by default 'gini'
    max_depth : int, optional
        树的最大深度, by default None
    random_state : int, optional
        随机种子, by default 42

    Returns
    -------
    DecisionTreeClassifier
        训练好的决策树模型
    """
    # 加载鸢尾花数据集
    iris = load_iris()
    X, y = iris.data, iris.target

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

    # 创建决策树模型
    clf = DecisionTreeClassifier(
        criterion=criterion, 
        max_depth=max_depth, 
        random_state=random_state
    )

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

    return clf

# 训练模型并评估
model = train_decision_tree(criterion='gini', max_depth=3)
print(f"Test accuracy: {model.score(X_test, y_test):.2f}")

# 可视化特征重要性
plt.barh(range(4), model.feature_importances_, align='center')
plt.yticks(range(4), load_iris().feature_names)
plt.xlabel('Feature Importance')
plt.ylabel('Features')
plt.title('Decision Tree Feature Importance')
plt.show()

5. 解决过拟合问题

决策树容易过拟合训练数据,可以通过以下方法缓解:

  1. 预剪枝 :在树完全生长之前停止分裂
  2. 限制最大深度 (max_depth)
  3. 设置叶节点最小样本数 (min_samples_leaf)
  4. 设置分裂最小样本数 (min_samples_split)

  5. 后剪枝 :先让树完全生长,然后删除不必要的子树

6. 处理类别不平衡

对于不平衡数据集,可以:

  1. 使用 class_weight 参数调整类别权重
  2. 对少数类样本进行过采样
  3. 对多数类样本进行欠采样
  4. 使用代价敏感学习

7. 超参数调优

使用 GridSearchCV 进行超参数搜索:

from sklearn.model_selection import GridSearchCV

param_grid = {'max_depth': [3, 5, 7, None],
    'min_samples_split': [2, 5, 10],
    'criterion': ['gini', 'entropy']
}

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

print(f"Best parameters: {grid_search.best_params_}")
print(f"Best cross-validation score: {grid_search.best_score_:.2f}")

8. 延伸思考

  1. 对于连续型特征,决策树如何选择最佳分裂点?
  2. 如何评估决策树模型的稳定性?
  3. 在什么情况下决策树会比其他分类器表现更好?

决策树作为一种直观易懂的机器学习算法,非常适合初学者入门。通过理解其基本原理和实现细节,可以为学习更复杂的集成方法如随机森林和梯度提升树打下坚实基础。

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