决策树实战:从Gini系数理解到Cart算法实现

1次阅读
没有评论

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

image.webp

认识 Gini 系数:决策树的纯度标尺

Gini 系数是 Cart 决策树用来衡量数据不纯度的核心指标。想象一个分类任务中,如果某个节点里所有样本都属于同一类别(比如全部是「合格产品」),那么这个节点就是完全纯净的,此时 Gini 系数为 0。计算公式如下:

决策树实战:从 Gini 系数理解到 Cart 算法实现

def gini_index(groups, classes):
    # groups: 分组后的样本集合
    # classes: 所有类别标签列表
    n_instances = sum(len(group) for group in groups)
    gini = 0.0
    for group in groups:
        size = len(group)
        if size == 0:
            continue
        score = 0.0
        for class_val in classes:
            p = [row[-1] for row in group].count(class_val) / size
            score += p * p
        gini += (1.0 - score) * (size / n_instances)
    return gini
  • 物理意义 :Gini 系数越小表示数据纯度越高,就像搅拌咖啡时奶油与咖啡混合越均匀,Gini 系数越大
  • 与信息增益对比
  • Gini 系数计算更快(无需对数运算)
  • 信息增益对类别分布更敏感
  • 实际效果差异通常小于 2%

Cart 算法实现四步走

  1. 数据准备阶段
    使用 sklearn 的鸢尾花数据集演示:

    from sklearn.datasets import load_iris
    iris = load_iris()
    X, y = iris.data, iris.target

  2. 特征划分评估
    遍历所有可能的划分点:

    def get_split(dataset):
        class_values = list(set(row[-1] for row in dataset))
        b_index, b_value, b_score, b_groups = 999, 999, 999, None
        for index in range(len(dataset[0])-1):
            for row in dataset:
                groups = test_split(index, row[index], dataset)
                gini = gini_index(groups, class_values)
                if gini < b_score:
                    b_index, b_value, b_score, b_groups = index, row[index], gini, groups
        return {"index":b_index, "value":b_value, "groups":b_groups}

  3. 递归构建决策树
    终止条件包括:

  4. 节点深度超过 max_depth
  5. 样本数小于 min_samples_split
  6. Gini 系数降为 0

  7. 预测与可视化
    使用 graphviz 输出决策树结构

过拟合的三大克星

  • 预剪枝 :提前停止树的生长
  • 设置 max_depth=3(典型值)
  • min_samples_leaf=5(每个叶节点最少样本)
  • 后剪枝 :CCP 代价复杂度剪枝
     from sklearn.tree import DecisionTreeClassifier
     clf = DecisionTreeClassifier(ccp_alpha=0.02)
  • 特征工程
  • 删除低方差特征
  • 使用 PCA 降维

项目实战经验分享

  1. 类别不平衡处理

    class_weight="balanced"  # 自动调整类别权重 

  2. 连续特征分箱技巧

    pd.cut(df["age"], bins=5)  # 等宽分箱 

  3. 树的可解释性优化

  4. 限制 max_features 参数
  5. 使用 feature_importances_筛选特征

思考进阶

  1. 当某个节点的 Gini 系数突然降为 0,可能暗示数据存在什么问题?
  2. 对于年龄、收入这类连续特征,Cart 算法是如何找到最佳分割点的?

(完整代码示例已上传 GitHub 仓库,包含数据集预处理和可视化模块)

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