决策树-学习笔记

决策树

决策树是树形结构的决策 / 机器学习模型,用分层分支做判断、分类、预测,像流程图一样直观:

根节点:整个判断起点

内部节点:判断条件(特征)

分支:条件的不同结果

叶节点:最终结论 / 分类结果

ID3决策树

ID3(Iterative Dichotomiser 3,迭代二分器 3)是最早的经典决策树算法。

分裂准则:信息增益 Information Gain

只能处理离散分类特征,不能直接处理连续值

只能做分类任务,不能回归

信息熵

条件熵

信息增益 = 信息熵 - 条件熵

计算过程

实例演算(ID3 计算全过程)

数据集(判断是否买电脑)

年龄 收入 学生 信用 购买电脑 (标签)
青年 一般
青年
中年 一般
老年 一般
老年 一般
老年
中年
青年 一般
青年 一般
老年 一般
青年
中年
中年 一般
老年

总样本 14 个:买 = 9,不买 = 5

步骤 1:计算总熵(信息熵) H (D)

步骤 2:分别计算 4 个特征信息增益

以年龄这列为例子:

青年条件熵:

中年条件熵:

老年条件熵:

此时条件熵(0.694) = 青年占比(5/14) * 青年条件熵(0.971) + 中年占比 (4/14)* 中年条件熵(0) + 老年占比(5/14) * 老年条件熵(0.971)

信息增益(信息熵 - 条件熵 = 信息增益):

同理算出其余特征增益:

收入增益 ≈ 0.029

学生增益 ≈ 0.151

信用增益 ≈ 0.048

步骤 3:选择最优特征(优先使用信息增益大的特征列,充当上层节点)

年龄信息增益最大,根节点 = 年龄,再对青年、中年、老年三个分支递归计算,直到叶子。

补充:

ID3 天然缺陷 ------ 信息增益偏向取值多的特征(种类多的特征)

如果特征列很混乱,就会导致信息增益大,这样算出来的信息增益没有意义,会影响判断节点顺序

C4.5决策树

信息增益率:

选择信息增益率最大的特征作为当前节点划分特征

信息增益率 = 信息增益 / 特征熵

特征熵的运算过程就是信息熵的运算过程,只不过是位置不同(一个统计标签分布,一个统计特征取值分布,相当于同一套熵运算)

CART决策树(分类)

CART决策树做分类操作时,优先采用基尼指数小的特征作为节点

基尼值

数据集 D 基尼:

Gini(D)=1−∑k=1Kpk2Gini(D)=1-\sum_{k=1}^K p_k^2Gini(D)=1−k=1∑Kpk2

pk:类别 k 样本占比

基尼越大 → 数据越混乱;纯样本基尼 = 0

若按某条件切分为 D1、D2,加权基尼:

Ginisplit=∣D1∣∣D∣Gini(D1)+∣D2∣∣D∣Gini(D2)Gini_{split} = \frac{|D_1|}{|D|}Gini(D_1) + \frac{|D_2|}{|D|}Gini(D_2)Ginisplit=∣D∣∣D1∣Gini(D1)+∣D∣∣D2∣Gini(D2)

遍历所有特征、所有分割点,选加权基尼最小的划分。

基尼指数

计算例子:

样本 年龄 (数值) 学生 (离散) 标签:购买电脑
1 20 0(不买)
2 22 0(不买)
3 35 1(买)
4 28 1(买)
5 40 1(买)
6 45 0(不买)

总样本 n=6;类别:是 = 3,否 = 3

步骤 1:计算原始数据集基尼

p是=3/6,p否=3/6 p_{是}=3/6, p_{否}=3/6 p是=3/6,p否=3/6

Gini(D)=1−((3/6)2+(3/6)2)=1−0.5=0.5 Gini(D)=1 - ( (3/6)^2+(3/6)^2)=1-0.5=0.5 Gini(D)=1−((3/6)2+(3/6)2)=1−0.5=0.5

步骤 2:遍历所有特征,计算每种分割的加权基尼

  1. 离散特征:学生(取值:是 / 否,只能二分)

    学生这列分为2个子集D1(不是学生)和D2(是学生)

    子集D1 = 学生 = 否:样本 1、2、3 → 买的 1个,不买的 2个

    Gini(D1)=1−(13)2−(23)2=1−1+49=49≈0.444 Gini(D_1)=1-\left(\frac13\right)^2-\left(\frac23\right)^2 = 1-\frac{1+4}{9}=\frac49≈0.444 Gini(D1)=1−(31)2−(32)2=1−91+4=94≈0.444

    子集D2 = 学生 = 是:样本 4、5、6 → 买的 2个,不买 1个

    Gini(D2)=1−(23)2−(13)2=49≈0.444 Gini(D_2)=1-\left(\frac23\right)^2-\left(\frac13\right)^2=\frac49≈0.444 Gini(D2)=1−(32)2−(31)2=94≈0.444

  2. 连续数值特征:年龄(找最优分割阈值)

    先将年龄进行排序:20, 22, 28, 35, 40, 45

    候选分割阈值:相邻两数取平均

    设排序后的数值数量 = n

    候选分割点总数 = n-1 ,所以有5个值

    t1=(20+22)/2=21t2=(22+28)/2=25t3=(28+35)/2=31.5t4=(35+40)/2=37.5t5=(40+45)/2=42.5 t1=(20+22)/2=21\\ t2=(22+28)/2=25\\ t3=(28+35)/2=31.5\\ t4=(35+40)/2=37.5\\ t5=(40+45)/2=42.5\\ t1=(20+22)/2=21t2=(22+28)/2=25t3=(28+35)/2=31.5t4=(35+40)/2=37.5t5=(40+45)/2=42.5

    1.分割阈值 t1=21(年龄 < 21 / ≥21)

    左边小于21岁的只有20岁一个,它的标签是0(不买)

    Gini左=1−(1/1)2=0 Gini_左=1-(1/1)^2=0 Gini左=1−(1/1)2=0

    右边大于等于21:22,28,35,40,45 (年龄集合)→ y(标签)=0,1,1,1,0,3 个 1(买),2 个 0(不买)

    p1=35,p0=25 p_1=\frac35,p_0=\frac25 p1=53,p0=52

    Gini右=1−(35)2−(25)2=1225=0.48Gini_右=1-(\frac{3}{5})^2-(\frac{2}{5})^2=\frac{12}{25}=0.48Gini右=1−(53)2−(52)2=2512=0.48

    加权基尼:

    1/6为左边的占比;5/6为右边的占比

    Gini21=16×0+56×0.48=0.40Gini_{21}=\frac16×0 + \frac56×0.48=0.40Gini21=61×0+65×0.48=0.40

    2.分割阈值 t=25(年龄 < 25 / ≥25)

    左 <25:20,22 → y=0,0

    Gini左=1−(2/2)2=0Gini_左=1 - (2/2)^2 = 0Gini左=1−(2/2)2=0

    右≥25:28,35,40,45 → y=1,1,1,0,3 个 1,1 个 0

    p1=34,p0=14p_1=\frac34,p_0=\frac14p1=43,p0=41

    Gini右=1−916−116=616=0.375Gini_右=1-\frac{9}{16}-\frac{1}{16}=\frac{6}{16}=0.375Gini右=1−169−161=166=0.375

    加权基尼:

    Gini25=26×0+46×0.375=0.25Gini_{25}=\frac{2}{6}×0+\frac{4}{6}×0.375=0.25Gini25=62×0+64×0.375=0.25

    3.分割阈值 t=31.5(年龄 < 31.5 / ≥31.5)

    左 <31.5:20,22,28 → y=0,0,1,1 个 1,2 个 0

    Gini左=49≈0.4444Gini_左=\frac49≈0.4444Gini左=94≈0.4444

    右≥31.5:35,40,45 → y=1,1,0,2 个 1,1 个 0

    Gini右=49≈0.4444Gini_右=\frac49≈0.4444Gini右=94≈0.4444

    加权基尼:

    Gini31.5=0.4444Gini_{31.5}=0.4444Gini31.5=0.4444

    4.分割阈值 t=37.5(年龄 < 37.5 / ≥37.5)

    左 <37.5:20,22,28,35 → y=0,0,1,1,2 个 1,2 个 0

    Gini左=1−0.52−0.52=0.5Gini_左=1-0.5^2-0.5^2=0.5Gini左=1−0.52−0.52=0.5

    右≥37.5:40,45 → y=1,0,1 个 1,1 个 0

    Gini右=0.5Gini_右=0.5Gini右=0.5

    加权基尼:

    Gini37.5=46×0.5+26×0.5=0.5Gini_{37.5}=\frac46×0.5+\frac26×0.5=0.5Gini37.5=64×0.5+62×0.5=0.5

    5.分割阈值 t=42.5(年龄 < 42.5 / ≥42.5)

    左 <42.5:20,22,28,35,40 → y=0,0,1,1,1,3 个 1,2 个 0

    Gini左=0.48Gini_左=0.48Gini左=0.48

    右≥42.5:45 → y=0

    Gini右=0Gini_右=0Gini右=0

    加权基尼:

    Gini42.5=56×0.48+16×0=0.40Gini_{42.5}=\frac56×0.48 + \frac16×0=0.40Gini42.5=65×0.48+61×0=0.40

    全部划分方案基尼汇总

    年龄 t=21:0.40

    年龄 t=25:0.25(全局最小值)

    年龄 t=31.5:0.4444

    年龄 t=37.5:0.5

    年龄 t=42.5:0.40

步骤 3:选最优划分

按学生划分:0.444

按年龄 25 划分:0.25(最小)

根节点:年龄 < 25

分支 1:年龄 <25 → 全是 "不买",叶子节点;

分支 2:年龄≥25,剩余 4 条样本,递归重复计算基尼,继续分裂。

CART决策树(回归)

节点样本标签均值 (\bar y)

MSE(D)=1n∑i=1n(yi−yˉ)2MSE(D)=\frac{1}{n}\sum_{i=1}^n(y_i-\bar y)^2MSE(D)=n1i=1∑n(yi−yˉ)2

分割后加权 MSE 最小为最优划分。

CART分类树和CART回归树区别

CART分类树预测输出的是一个离散值,CART回归树预测输出的是一个连续值

CART分类树使用基尼指数作为划分、构建树的依据,CART回归树使用平方损失

分类树使用叶子节点多数类别作为预测类别,回归树则采用叶子节点里的均值作为预测输出

计算例子:

数据集:房屋面积 (x)→房价 y(连续数值标签)

面积 x 房价 y (万)
60 100
70 120
90 180
110 220
130 260

目标:找面积最优分割点,最小化加权 MSE

整体均值: yˉ=(100+120+180+220+260)/5=176\bar y=(100+120+180+220+260)/5=176yˉ=(100+120+180+220+260)/5=176

整体 MSE:

MSE=(100−176)2+(120−176)2+(180−176)2+(220−176)2+(260−176)25=3406.4MSE= \frac{(100-176)^2+(120-176)^2+(180-176)^2+(220-176)^2+(260-176)^2}{5}=3406.4MSE=5(100−176)2+(120−176)2+(180−176)2+(220−176)2+(260−176)2=3406.4

候选分割点:65、80、100、120

举例:以分割阈值 80(x<80 /x≥80)

D1 (x<80):60,70 → y=100,120,均值 = 110

MSE1=(100−110)2+(120−110)22=100MSE_1=\frac{(100-110)^2+(120-110)^2}{2}=100MSE1=2(100−110)2+(120−110)2=100

D2 (x≥80):90,110,130 → y=180,220,260,均值 = 220

MSE2=(180−220)2+(220−220)2+(260−220)23≈1066.67MSE_2=\frac{(180-220)^2+(220-220)^2+(260-220)^2}{3}≈1066.67MSE2=3(180−220)2+(220−220)2+(260−220)2≈1066.67

加权 MSE:

25×100+35×1066.67≈680\frac{2}{5}×100 + \frac{3}{5}×1066.67≈68052×100+53×1066.67≈680

遍历所有阈值后,该分割加权 MSE 最小,作为第一层分裂。

叶子节点预测值 = 当前子集 y 的平均值。

决策树正则化(决策树剪枝)

决策树剪枝就是一种防止决策树过拟合的一种正则化方法,提高其泛化能力

剪枝:把子树的节点全部删掉,使用叶子节点来替换

  1. 预剪枝(先限制,边建边剪,简单高效)

    在构建决策树递归分裂的过程中,提前设置停止条件,不允许继续分裂,从源头避免树长得太深。

  2. 后剪枝(先生完整大树,再反向修剪,效果更好)

    流程:先不设限制,训练一棵完全生长、充分分裂的完整树,再从底层叶子往根节点反向遍历,判断是否要剪掉某棵子树。

剪枝操作:

py 复制代码
from sklearn.tree import DecisionTreeClassifier
# 预剪枝:设置深度、最小样本限制
tree_pre = DecisionTreeClassifier(max_depth=4, min_samples_split=5)

# CART代价复杂度后剪枝
tree_post = DecisionTreeClassifier(ccp_alpha=0.01) # alpha越大剪得越狠

总结

ID3树的分支方式:信息增益

C4.5树的分支方式:信息增益率

cart树的分支方式:基尼指数

相关推荐
MartinYeung51 小时前
[论文学习]WASP:面向提示注入攻击的Web代理安全性基准测试
前端·网络·学习
懒狗跑ai的程序员Brain1 小时前
如何利用粒子系统在unity制作一个烟花(粒子系统学习向)
学习·unity·游戏引擎
念何架构之路1 小时前
moby-BuildKit(builder-next)
学习·docker·容器
hongmai6668882 小时前
ESP32-C5-WROOM-1-N16R8:双频Wi-Fi 6与多协议融合,重新定义物联网连接新标准
笔记·嵌入式硬件·物联网·智能路由器·risc-v
风曦Kisaki2 小时前
Kubernetes(K8s)笔记Day04:控制器(ReplicaSet 与Deployment),滚动更新及回滚,滚动更新策略,Pod 的 DNS 策略
linux·运维·笔记·docker·容器·kubernetes
惠惠软件2 小时前
快速生成磁盘目录文件-双击运行a.bat ,可以看到磁盘目录下出现了 文件目录.txt-供大家学习研究参考
windows·学习·文件列表
星恒随风2 小时前
C++ STL 详解:set 与 multiset 的使用、区间查询和算法应用
开发语言·c++·笔记·学习·算法
m4Rk_2 小时前
【论文阅读】Agent 记忆机制(20):RecMem——只在信息反复出现时进行长期记忆巩固
论文阅读·人工智能·学习·开源·github
菩提树下的打坐2 小时前
测试工程与 DevOps / SRE 的边界:三方协作的真实工作流
学习