决策树
决策树是树形结构的决策 / 机器学习模型,用分层分支做判断、分类、预测,像流程图一样直观:
根节点:整个判断起点
内部节点:判断条件(特征)
分支:条件的不同结果
叶节点:最终结论 / 分类结果
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:遍历所有特征,计算每种分割的加权基尼
-
离散特征:学生(取值:是 / 否,只能二分)
学生这列分为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
-
连续数值特征:年龄(找最优分割阈值)
先将年龄进行排序: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 的平均值。
决策树正则化(决策树剪枝)
决策树剪枝就是一种防止决策树过拟合的一种正则化方法,提高其泛化能力
剪枝:把子树的节点全部删掉,使用叶子节点来替换
-
预剪枝(先限制,边建边剪,简单高效)
在构建决策树递归分裂的过程中,提前设置停止条件,不允许继续分裂,从源头避免树长得太深。
-
后剪枝(先生完整大树,再反向修剪,效果更好)
流程:先不设限制,训练一棵完全生长、充分分裂的完整树,再从底层叶子往根节点反向遍历,判断是否要剪掉某棵子树。
剪枝操作:
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树的分支方式:基尼指数