决策树算法实战指南:3种核心模型对比与新手避坑手册

1次阅读
没有评论

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

image.webp

为什么需要决策树?

决策树是一种模拟人类决策过程的算法,在实际业务中应用广泛。比如在金融风控中,银行需要根据用户的年龄、收入、信用历史等特征,快速判断是否批准贷款申请;在医疗诊断中,医生会根据病人的症状、检查结果等特征,判断患某种疾病的可能性。决策树能够清晰地展示决策路径,让非技术人员也能理解模型的判断逻辑。

决策树算法实战指南:3 种核心模型对比与新手避坑手册

三种经典决策树算法对比

1. ID3 算法:基于信息增益

ID3 算法使用信息增益来选择最优分裂特征。信息增益表示分裂前后信息不确定性的减少程度,计算公式为:

Gain(D,A) = Entropy(D) - ∑(|D_v|/|D|)*Entropy(D_v)

其中,Entropy(D) = -∑p_i*log2(p_i) 表示数据集 D 的信息熵。

ID3 的缺点是倾向于选择取值较多的特征,可能导致过拟合。

2. C4.5 算法:基于增益率

C4.5 算法改进了 ID3,使用增益率来减少对多值特征的偏好:

Gain_ratio(D,A) = Gain(D,A)/SplitInfo(D,A)

其中,SplitInfo(D,A) = -∑(|D_v|/|D|)*log2(|D_v|/|D|) 称为分裂信息。

3. CART 算法:基于基尼系数

CART 算法可以用于分类和回归任务,分类时使用基尼系数:

Gini(D) = 1 - ∑p_i^2

基尼系数越小,表示数据集的纯度越高。

代码实战:sklearn 实现三种决策树

环境准备

import numpy as np
from sklearn.tree import DecisionTreeClassifier, export_text, plot_tree
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
import matplotlib.pyplot as plt

数据准备

# 加载鸢尾花数据集
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.2, random_state=42)

ID3 风格决策树

虽然 sklearn 没有直接实现 ID3,但可以通过设置 criterion=’entropy’ 来模拟:

tree_id3 = DecisionTreeClassifier(criterion='entropy', max_depth=3)
tree_id3.fit(X_train, y_train)
print(export_text(tree_id3, feature_names=iris.feature_names))

C4.5 风格决策树

sklearn 的 DecisionTreeClassifier 默认使用改进后的信息增益,可以近似 C4.5:

tree_c45 = DecisionTreeClassifier(criterion='entropy', max_depth=3, min_samples_split=5)
tree_c45.fit(X_train, y_train)

CART 决策树

tree_cart = DecisionTreeClassifier(criterion='gini', max_depth=3)
tree_cart.fit(X_train, y_train)

可视化决策树

plt.figure(figsize=(12,8))
plot_tree(tree_cart, feature_names=iris.feature_names, class_names=iris.target_names, filled=True)
plt.show()

专项问题分析

解决过拟合问题

  1. 预剪枝 :通过设置参数提前停止树的生长
  2. max_depth:树的最大深度
  3. min_samples_split:节点分裂所需的最小样本数
  4. min_samples_leaf:叶节点所需的最小样本数

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

处理类别不平衡

# 设置 class_weight 参数
tree = DecisionTreeClassifier(class_weight='balanced')

或者手动指定权重:

class_weight = {0:1, 1:5, 2:1}  # 给类别 1 更高的权重
tree = DecisionTreeClassifier(class_weight=class_weight)

生产环境建议

  1. 特征工程
  2. 决策树对单调变换不敏感,无需标准化
  3. 但对缺失值敏感,需要提前处理
  4. 离散特征更好,可以考虑分箱连续特征

  5. 模型集成

  6. 随机森林:多棵决策树的集成
  7. GBDT:梯度提升决策树
  8. XGBoost/LightGBM:优化的 GBDT 实现

开放性问题

  1. 决策树本质上是通过轴平行分割来处理非线性数据,对于复杂的非线性关系,可能需要很深的树或者考虑其他模型。
  2. 在实时推理场景下,可以通过限制树的深度、使用更简单的树结构,或者将决策规则提前编译成 if-else 语句来优化性能。

决策树是机器学习入门的重要算法,理解它的原理和实现是打好基础的关键。希望通过这篇文章,你能对三种主要决策树算法有清晰的认识,并在实际项目中合理应用。

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