共计 1690 个字符,预计需要花费 5 分钟才能阅读完成。
背景介绍
决策树是一种经典的机器学习算法,它通过树形结构对数据进行分类或回归。CART(Classification and Regression Trees)是决策树的一种实现方式,既可以用于分类任务,也可以用于回归任务。CART 决策树的特点包括:

- 使用基尼不纯度(Gini impurity)作为分类任务的分裂标准
- 二叉树结构,每个节点最多有两个子节点
- 能够处理数值型和类别型特征
- 对异常值相对鲁棒
案例解析:鸢尾花分类
我们以经典的鸢尾花数据集为例,演示 CART 决策树的构建过程。这个数据集包含 150 个样本,每个样本有 4 个特征(花萼长度、花萼宽度、花瓣长度、花瓣宽度)和 1 个类别标签(山鸢尾、变色鸢尾、维吉尼亚鸢尾)。
决策树的构建过程大致如下:
- 从根节点开始,选择最优特征和分割点
- 根据选择的分割点将数据集分成两个子集
- 对每个子集递归地重复上述过程,直到满足停止条件
- 生成最终的决策树
代码实现
下面是使用 scikit-learn 实现 CART 决策树的完整代码示例:
from sklearn.datasets import load_iris
from sklearn.tree import DecisionTreeClassifier, export_text
# 加载数据
iris = load_iris()
X = iris.data
y = iris.target
# 创建决策树分类器
tree_clf = DecisionTreeClassifier(criterion='gini', max_depth=3, random_state=42)
# 训练模型
tree_clf.fit(X, y)
# 打印决策树规则
tree_rules = export_text(tree_clf, feature_names=iris.feature_names)
print("决策树规则:")
print(tree_rules)
代码说明:
criterion='gini'表示使用基尼不纯度作为分裂标准max_depth=3限制树的最大深度为 3random_state=42确保每次运行结果一致
参数调优
CART 决策树有几个关键参数需要调优:
max_depth:树的最大深度。值越大模型越复杂,可能过拟合;值太小可能欠拟合。min_samples_split:节点分裂所需的最小样本数。可以防止过拟合。min_samples_leaf:叶节点所需的最小样本数。max_features:寻找最佳分割时考虑的特征数量。
调优建议:
- 开始时使用较小的
max_depth(如 3 -5),然后逐步增加 - 使用交叉验证来评估不同参数组合的效果
- 考虑使用
GridSearchCV进行自动化参数搜索
避坑指南
在使用 CART 决策树时,常见的误区包括:
- 忽视特征缩放:虽然决策树对特征缩放不敏感,但当特征尺度差异很大时可能影响分裂点选择。
- 过拟合:决策树容易过拟合训练数据,需要通过参数调优和剪枝来避免。
- 忽略类别不平衡:对于类别不平衡的数据,可以考虑设置
class_weight参数。 - 忽视特征重要性:训练完成后可以通过
feature_importances_属性查看特征重要性。
性能考量
CART 决策树的性能特点:
- 时间复杂度:训练时间复杂度为 O(n_features * n_samples^2 * log(n_samples))
- 空间复杂度:存储整个树需要 O(n_nodes)的空间
- 预测速度快:预测时间复杂度为 O(tree_depth)
实际应用建议:
- 对于大型数据集,考虑设置
max_depth限制树的大小 - 如果内存有限,可以使用较小的
min_samples_leaf和min_samples_split - 考虑使用增量学习(partial_fit)处理非常大的数据集
总结
通过这个简单的鸢尾花分类案例,我们了解了 CART 决策树的原理和实现方法。决策树是一种直观且强大的机器学习算法,特别适合作为入门学习的第一个算法。虽然它有一些局限性,如容易过拟合,但通过合理的参数调优和剪枝,可以获得很好的效果。
在实际项目中,建议先使用决策树建立基线模型,然后再尝试更复杂的算法。决策树也可以作为特征选择的工具,通过分析特征重要性来指导后续建模。
正文完
发表至: 机器学习
近三天内
