决策树模型实战:从零推导CART算法中的Gini系数计算

1次阅读
没有评论

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

image.webp

1. 决策树与 Gini 系数基础

决策树 (Decision Tree) 是一种通过树状结构做决策的机器学习算法。CART(Classification and Regression Trees)算法使用 Gini 系数 (Gini Index) 作为特征选择的标准,它衡量数据的不纯度。

决策树模型实战:从零推导 CART 算法中的 Gini 系数计算

  • Gini 系数直观理解:可以看作是从数据集中随机抽取两个样本,它们属于不同类别的概率。值越小表示数据越纯。
  • 为什么要用 Gini 系数:计算简单且不需要对数运算,适合大规模数据。

2. Gini 系数计算公式推导

Gini 系数的数学定义为:

Gini = 1 - Σ(p_i)^2

其中 p_i 是第 i 类样本所占比例。

举个例子,假设一个节点有 10 个样本:

  • 情况 A:5 个类别 0,5 个类别 1
    Gini = 1 – (0.5² + 0.5²) = 0.5

  • 情况 B:9 个类别 0,1 个类别 1
    Gini = 1 – (0.9² + 0.1²) = 0.18

可以看出情况 B 的 Gini 值更小,纯度更高。

3. Python 代码实现

import numpy as np

def compute_gini(y):
    """
    计算 Gini 系数
    :param y: 目标值数组
    :return: Gini 系数值
    """
    if len(y) == 0:
        return 0
    # 计算每个类别的比例
    proportions = np.bincount(y) / len(y)
    return 1 - np.sum(proportions ** 2)

# 示例用法
y = np.array([0, 0, 1, 1, 1])  # 2 个类别 0,3 个类别 1
print(f"Gini 系数: {compute_gini(y):.3f}")  # 输出: 0.480

4. 与信息增益的对比

  • 计算效率:Gini 系数计算更快(无对数运算)
  • 结果倾向:Gini 倾向于分裂出更大的分区
  • 适用场景
  • Gini:默认选择,计算高效
  • 信息增益:需要更平衡的树时

5. sklearn 实战示例

from sklearn.tree import DecisionTreeClassifier
from sklearn.datasets import load_iris

# 加载数据
iris = load_iris()
X, y = iris.data, iris.target

# 使用 gini 准则创建决策树
clf = DecisionTreeClassifier(criterion='gini', max_depth=3)
clf.fit(X, y)

# 查看特征重要性
print("特征重要性:", clf.feature_importances_)

6. 常见问题 FAQ

Q1: 如何处理连续特征?

A: CART 会对连续特征排序后尝试所有可能的分割点,选择 Gini 系数最小的分割方式。

Q2: 类别不平衡时 Gini 系数会怎样?

A: Gini 系数对多数类更敏感。可以尝试加权 Gini 或改用其他指标。

7. 延伸思考

  1. Gini 系数的取值范围是[0,1],0 表示完全纯净,1 表示最大不纯度
  2. CART 采用二分法因为计算更高效,且可以处理多分类问题
  3. 当所有样本属于同一类时 Gini 值为 0(完全纯净)

总结

本文从基础概念到代码实现,详细讲解了 CART 决策树中 Gini 系数的计算原理。理解这一核心指标对于掌握决策树算法至关重要。建议读者动手实现代码,在实际数据上观察 Gini 系数的变化规律。

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