机器学习基础:决策树
引言:
决策树模型是一个基于树的一个监督学习算法,采用递归的方式进行分割数据,最终得到一颗树,根据这棵树的叶子节点我们就可以预测结果,从而达到分类的目的。本篇文章将采用鸢尾花数据集进行分析,并使用scikit-learn和鸢尾花进行简单模型的构建和评估。
原理介绍:
决策树是一个递归生成过程的结果,面对一堆具有不同属性数据,我们根据不同类别样本具有不同的属性进行区分,重复进行这一步骤。我们有数据集{ x 1 , x 2 , x 3 . . . x n x_1,x_2,x_3...x_n x1,x2,x3...xn}和属性集{ a 1 , a 2 , a 3 . . . a n a_1,a_2,a_3...a_n a1,a2,a3...an}当做到以下三种情况的时候,递归停止:
(1)当前节点所包含的样本都属于同一类别。
(2)当属性集为空或者所有样本在属性集的取值都相同时。
(3)当前节点所包含的样本集为空时。
这样我们就可以将数据分成了一个树的结构,这颗树的叶子节点呈现出以下的特点:
(1)节点中的样本都为同一类别
(2)节点中的样本属性取值都相同
(3)节点为空节点
对于这些特点(2)和(3)都是不允许的,我们需要认为设定一些规则,对于特点(2)我们将类别设定为节点中样本数量最多的类别,对于(3)我们设定该节点的类别为其父节点样本数量最多的类别。
信息熵
信息熵是度量样本集合纯度最常用的一种指标,假定当前样本集合D中第k类样本所占的比列为
p
k
(
k
=
1
,
2
,
.
.
.
,
∣
y
∣
)
p_k(k = 1,2,...,|y|)
pk(k=1,2,...,∣y∣),则D的信息熵定义为:
E
n
t
(
D
)
=
−
∑
k
=
1
∣
y
∣
p
k
l
o
g
2
p
k
Ent(D) = -\sum_{k = 1}^{|y|}p_klog_2p_k
Ent(D)=−k=1∑∣y∣pklog2pk
信息增益
假定离散属性a有V个可能的取值
(
a
1
,
a
2
,
a
3
,
.
.
.
,
a
V
)
(a^1,a^2,a^3,...,a^V)
(a1,a2,a3,...,aV),若采用a来对样本集合D进行划分就会出现V个不同的节点,由此我们可以通过上式信息熵的计算公式来得出信息增益的计算公式:
G
a
i
n
(
D
,
a
)
=
E
n
t
(
D
)
−
∑
v
=
1
V
∣
D
v
∣
∣
D
∣
E
n
t
(
D
v
)
Gain(D,a) = Ent(D)-\sum_{v = 1}^{V}\frac{|D^v|}{|D|}Ent(D^v)
Gain(D,a)=Ent(D)−v=1∑V∣D∣∣Dv∣Ent(Dv)
因为属性相同的样本数量占有D的比例不同,不同的子节点我们采用不同的权重
∣
D
v
∣
∣
D
∣
\frac{|D^v|}{|D|}
∣D∣∣Dv∣来进行表示。
由此我们可以利用鸢尾花数据集进行信息熵和信息增益的一个计算演示,对于著名的鸢尾花数据集来说,他有四个属性值花萼长度(Sepal Length),花萼宽度(Sepal Width),花瓣长度(Petal Length),花瓣宽度(Petal Width),三个类别(Setosa、Versicolour、Virginica)一共150组数据样本,对于鸢尾花数据集我们可以通过以下代码进行查看。
from sklearn.datasets import load_iris
load = load_iris()
print(load)
或者通过https://www.modelscope.cn/datasets/alvisx/iris1/file/view/master/ir.csv进行下载。
首先我们对全体样本集进行信息熵的计算:
E
n
t
(
D
)
=
−
∑
k
=
1
3
p
k
l
o
g
2
p
k
=
−
(
1
3
l
o
g
2
1
3
+
1
3
l
o
g
2
1
3
+
1
3
l
o
g
2
1
3
)
=
l
o
g
2
3
Ent(D) = -\sum_{k = 1}^{3}p_klog_2p_k=-(\frac{1}{3}log_2\frac{1}{3}+\frac{1}{3}log_2\frac{1}{3}+\frac{1}{3}log_2\frac{1}{3}) = log_23
Ent(D)=−k=1∑3pklog2pk=−(31log231+31log231+31log231)=log23
假设我们依据属性花萼长度(Sepal Length)对样本进行划分,可以划分35个节点,有些节点可能只有一个样本,所以只有一个样本的节点它的信息熵就是:
E
n
t
(
D
v
)
=
−
∑
k
=
1
1
1
l
o
g
2
1
=
0
Ent(D_v) = -\sum_{k = 1}^11log_21 = 0
Ent(Dv)=−k=1∑11log21=0
所有的样本情况是:
`{'5.1': 9, '4.9': 6, '4.7': 2, '4.6': 4, '5.0': 10, '5.4': 6, '4.4': 3, '4.8': 5, '4.3': 1, '5.8': 7, '5.7': 8, '5.2': 4, '5.5': 7, '4.5': 1, '5.3': 1, '7.0': 1, '6.4': 7, '6.9': 4, '6.5': 5, '6.3': 9, '6.6': 2, '5.9': 3, '6.0': 6, '6.1': 6, '5.6': 6, '6.7': 8, '6.2': 4, '6.8': 3, '7.1': 1, '7.6': 1, '7.3': 1, '7.2': 3, '7.7': 4, '7.4': 1, '7.9': 1}
通过一些处理,我们就可以得到:
{'5.1': ['0', '0', '0', '0', '0', '0', '0', '0', '1'], '4.9': ['0', '0', '0', '0', '1', '2'], '4.7': ['0', '0'], '4.6': ['0', '0', '0', '0'], '5.0': ['0', '0', '0', '0', '0', '0', '0', '0', '1', '1'], '5.4': ['0', '0', '0', '0', '0', '1'], '4.4': ['0', '0', '0'], '4.8': ['0', '0', '0', '0', '0'], '4.3': ['0'], '5.8': ['0', '1', '1', '1', '2', '2', '2'], '5.7': ['0', '0', '1', '1', '1', '1', '1', '2'], '5.2': ['0', '0', '0', '1'], '5.5': ['0', '0', '1', '1', '1', '1', '1'], '4.5': ['0'], '5.3': ['0'], '7.0': ['1'], '6.4': ['1', '1', '2', '2', '2', '2', '2'], '6.9': ['1', '2', '2', '2'], '6.5': ['1', '2', '2', '2', '2'], '6.3': ['1', '1', '1', '2', '2', '2', '2', '2', '2'], '6.6': ['1', '1'], '5.9': ['1', '1', '2'], '6.0': ['1', '1', '1', '1', '2', '2'], '6.1': ['1', '1', '1', '1', '2', '2'], '5.6': ['1', '1', '1', '1', '1', '2'], '6.7': ['1', '1', '1', '2', '2', '2', '2', '2'], '6.2': ['1', '1', '2', '2'], '6.8': ['1', '2', '2'], '7.1': ['2'], '7.6': ['2'], '7.3': ['2'], '7.2': ['2', '2', '2'], '7.7': ['2', '2', '2', '2'], '7.4': ['2'], '7.9': ['2']}
对于这些样本划分好的节点,我们可以计算其每个节点的信息熵,对于第一个节点有:
E
n
t
(
D
1
)
=
−
∑
k
=
1
2
p
k
l
o
g
2
p
k
=
−
(
8
9
l
o
g
2
8
9
+
1
9
l
o
g
2
1
9
)
=
0.50325
Ent(D^1) = -\sum_{k = 1}^2 p_klog_2p_k = -(\frac{8}{9}log_2\frac{8}{9}+\frac{1}{9}log_2\frac{1}{9}) = 0.50325
Ent(D1)=−k=1∑2pklog2pk=−(98log298+91log291)=0.50325
我们可以使用deepseekAI继续帮我们处理接下来的34个点:
{'5.1': 0.5032583347756457, '4.9': 1.2516291673878228, '4.7': -0.0, '4.6': -0.0, '5.0': 0.7219280948873623, '5.4': 0.6500224216483541, '4.4': -0.0, '4.8': -0.0, '4.3': -0.0, '5.8': 1.4488156357251847, '5.7': 1.2987949406953985, '5.2': 0.8112781244591328, '5.5': 0.863120568566631, '4.5': -0.0, '5.3': -0.0, '7.0': -0.0, '6.4': 0.863120568566631, '6.9': 0.8112781244591328, '6.5': 0.7219280948873623, '6.3': 0.9182958340544896, '6.6': -0.0, '5.9': 0.9182958340544896, '6.0': 0.9182958340544896, '6.1': 0.9182958340544896, '5.6': 0.6500224216483541, '6.7': 0.954434002924965, '6.2': 1.0, '6.8': 0.9182958340544896, '7.1': -0.0, '7.6': -0.0, '7.3': -0.0, '7.2': -0.0, '7.7': -0.0, '7.4': -0.0, '7.9': -0.0}
这样35个点的信息熵就都算出来了,接下来我们算信息增益,继续通过AI进行辅助求求解可以得到:
G
a
i
n
(
D
,
花萼长度
)
=
E
n
t
(
D
)
−
∑
v
=
1
35
∣
D
v
∣
∣
D
∣
E
n
t
(
D
v
)
=
l
o
g
2
3
−
0.7080248798300983
=
0.87694
Gain(D,花萼长度) = Ent(D)-\sum_{v = 1}^{35}\frac{|D^v|}{|D|}Ent(D^v) = log_23 - 0.7080248798300983 = 0.87694
Gain(D,花萼长度)=Ent(D)−v=1∑35∣D∣∣Dv∣Ent(Dv)=log23−0.7080248798300983=0.87694
同理,我们可以得到:
G
a
i
n
(
D
,
花萼宽度
)
=
0.5167
,
G
a
i
n
(
D
,
花瓣长度
)
=
1.4463
Gain(D,花萼宽度) = 0.5167 ,Gain(D,花瓣长度) = 1.4463
Gain(D,花萼宽度)=0.5167,Gain(D,花瓣长度)=1.4463
G
a
i
n
(
D
,
花瓣宽度
)
=
1.4359
,
G
a
i
n
(
D
,
花萼长度
)
=
0.8769
Gain(D,花瓣宽度) = 1.4359,Gain(D,花萼长度) = 0.8769
Gain(D,花瓣宽度)=1.4359,Gain(D,花萼长度)=0.8769
因为花瓣长度获得的信息增益最大,由此我们可以基于"花瓣长度"进行划分。
决策树构建过程演示
在上一步中,我们已经计算出四个属性中的信息增益,并确定花瓣长度(Petal Length)具有最大的信息增益值(1.4463),因此它将成为决策树的根节点,用于对整个数据集D进行第一次分裂。接下来,我们将基于花瓣长度对数据集进行划分,并递归地重复上述过程,直到满足停止条件。这将帮助我们逐步构建出一棵完整的决策树。
为了演示,我们首先需要查看花瓣长度的分布情况。鸢尾花数据集中的花瓣长度取值范围较广(从1.0到6.9),共有16个不同的离散值(假设我们将连续值离散化为这些具体测量值)。通过类似的前述代码处理,我们可以得到每个花瓣长度值的样本计数和类别标签分布:
from collections import Counter
import numpy as np
from sklearn.datasets import load_iris
# 加载数据集
iris = load_iris()
X = iris.data[:, 2] # 提取花瓣长度(第3列,索引2)
y = iris.target # 类别标签:0=Setosa, 1=Versicolour, 2=Virginica
# 统计每个花瓣长度值的样本计数和标签分布
petal_length_groups = {}
for length in np.unique(X):
mask = X == length
counts = Counter(y[mask])
labels = [str(label) for label in y[mask]] # 转换为字符串便于显示
petal_length_groups[str(length)] = {
'count': len(labels),
'labels': labels,
'class_counts': dict(counts)
}
print(petal_length_groups)
运行后,我们得到类似以下的分布(简化显示,仅列出关键信息):
{'1.0': {'count': 7, 'labels': ['0']*7, 'class_counts': {0: 7}},
'1.1': {'count': 7, 'labels': ['0']*7, 'class_counts': {0: 7}},
'1.2': {'count': 4, 'labels': ['0']*4, 'class_counts': {0: 4}},
'1.3': {'count': 6, 'labels': ['0']*6, 'class_counts': {0: 6}},
'1.4': {'count': 5, 'labels': ['0']*5, 'class_counts': {0: 5}},
'1.5': {'count': 3, 'labels': ['0']*3, 'class_counts': {0: 3}},
'1.6': {'count': 1, 'labels': ['0'], 'class_counts': {0: 1}},
'1.7': {'count': 4, 'labels': ['0']*4, 'class_counts': {0: 4}},
'3.0': {'count': 8, 'labels': ['1']*8, 'class_counts': {1: 8}},
'3.3': {'count': 1, 'labels': ['1'], 'class_counts': {1: 1}},
'3.5': {'count': 2, 'labels': ['1']*2, 'class_counts': {1: 2}},
'4.0': {'count': 2, 'labels': ['1', '2'], 'class_counts': {1: 1, 2: 1}},
'4.4': {'count': 3, 'labels': ['1', '2', '2'], 'class_counts': {1: 1, 2: 2}},
'4.5': {'count': 1, 'labels': ['2'], 'class_counts': {2: 1}},
'4.8': {'count': 5, 'labels': ['2']*5, 'class_counts': {2: 5}},
'4.9': {'count': 2, 'labels': ['2']*2, 'class_counts': {2: 2}},
'5.0': {'count': 1, 'labels': ['2'], 'class_counts': {2: 1}},
'5.1': {'count': 1, 'labels': ['2'], 'class_counts': {2: 1}},
'5.2': {'count': 2, 'labels': ['2']*2, 'class_counts': {2: 2}},
'5.4': {'count': 3, 'labels': ['2']*3, 'class_counts': {2: 3}},
'5.5': {'count': 1, 'labels': ['2'], 'class_counts': {2: 1}},
'5.7': {'count': 2, 'labels': ['2']*2, 'class_counts': {2: 2}},
'5.8': {'count': 1, 'labels': ['2'], 'class_counts': {2: 1}},
'5.9': {'count': 1, 'labels': ['2'], 'class_counts': {2: 1}},
'6.0': {'count': 1, 'labels': ['2'], 'class_counts': {2: 1}},
'6.1': {'count': 1, 'labels': ['2'], 'class_counts': {2: 1}},
'6.3': {'count': 1, 'labels': ['2'], 'class_counts': {2: 1}},
'6.7': {'count': 1, 'labels': ['2'], 'class_counts': {2: 1}}}
从分布中可以看出,花瓣长度小于等于2.0的样本全部属于Setosa类(类别0),这是一个纯节点(Ent=0),无需进一步分裂。花瓣长度在3.0到3.5之间的样本全部属于Versicolour类(类别1),同样纯净。剩余的较大值(4.0以上)主要属于Virginica类(类别2),但有一些混杂(如4.0和4.4有少量Versicolour)。这产生了三个主要子集:D1(小花瓣,Setosa)、D2(中花瓣,Versicolour)和D3(大花瓣,混合Versicolour和Virginica)。
对于子集D3,我们需要递归计算剩余属性的信息增益。假设我们对D3(约50个样本)计算四个属性的Gain,发现花瓣宽度(Petal Width)的信息增益最高(约0.95),因此在D3的根节点使用花瓣宽度进行分裂。类似地,继续递归,直到所有叶子节点满足停止条件。
通过这个过程,我们最终得到一棵决策树,大致结构如下(简化表示):
根节点:花瓣长度 ≤ 2.45? 是 → Setosa(纯节点,50个样本) 否 → 子节点1:花瓣长度 ≤ 4.95? 是 →
Versicolour(纯节点,54个样本) 否 → 子节点2:花瓣宽度 ≤ 1.75? 是 → Versicolour(纯节点,约几样本) 否 → Virginica(纯节点,46个样本)
这棵树非常浅(深度约3),因为鸢尾花数据集特征较少且类别分离明显。在实际构建中,我们可以设置参数如最大深度(max_depth=3)或最小样本分裂数(min_samples_split=2)来控制树的复杂度,避免过拟合。
使用scikit-learn实现决策树
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
from sklearn.tree import DecisionTreeClassifier, plot_tree
from sklearn.metrics import accuracy_score, confusion_matrix, classification_report
import matplotlib.pyplot as plt
import numpy as np
# 加载数据集
iris = load_iris()
X, y = iris.data, iris.target
# 分割数据集:80%训练,20%测试
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42, stratify=y)
# 构建决策树模型,使用信息熵准则
clf = DecisionTreeClassifier(criterion='entropy', max_depth=3, random_state=42)
clf.fit(X_train, y_train)
# 预测
y_pred = clf.predict(X_test)
# 评估准确率
accuracy = accuracy_score(y_test, y_pred)
print(f"模型准确率: {accuracy:.4f}")
# 混淆矩阵
cm = confusion_matrix(y_test, y_pred)
print("混淆矩阵:\n", cm)
# 分类报告
print("\n分类报告:\n", classification_report(y_test, y_pred, target_names=iris.target_names))
# 可视化决策树(可选,需安装graphviz)
plt.figure(figsize=(12, 8))
plot_tree(clf, feature_names=iris.feature_names, class_names=iris.target_names, filled=True)
plt.title("鸢尾花决策树可视化")
plt.show()
运行结果示例(基于随机种子42):

从结果可见,在测试集上准确率达到100%,这得益于鸢尾花数据集的线性可分性。混淆矩阵显示无误分类,分类报告中的精确率、召回率和F1分数均为1.0。决策树可视化图直观展示了分裂路径,与我们手动计算的结构高度一致(根节点为花瓣长度)。
模型评估与优缺点分析
评估指标
除了准确率,我们还可以使用其他指标评估模型:
精确率(Precision):预测为正类的样本中真正例的比例。
召回率(Recall):实际正类中被正确预测的比例。
F1分数:精确率和召回率的调和平均,适用于不平衡数据集。
在鸢尾花上,由于类别均衡,这些指标表现优异。但在实际应用中,应使用交叉验证(e.g., cross_val_score)来更稳健地评估。
决策树的优缺点
优点:
可解释性强:树结构易于理解和可视化。
无需数据预处理:能处理数值和类别特征。
非参数化:不假设数据分布。
缺点:
易过拟合:树过深时需剪枝(pre-pruning或post-pruning)。
对噪声敏感:小变化可能导致树结构大变。
偏向高基数特征:信息增益可能偏好多值属性。
为缓解过拟合,可调整min_samples_leaf或使用集成方法如随机森林。
结论
通过本文,我们从理论到实践全面探讨了决策树算法,使用鸢尾花数据集演示了信息熵和信息增益的计算过程,并借助scikit-learn快速构建了高效模型。决策树作为监督学习的基础工具,在分类任务中发挥关键作用,尤其适合初学者入门。未来,可扩展到回归树或与其他算法结合(如XGBoost)以提升性能。
如果您对代码有疑问或想深入某个部分,欢迎留言讨论!数据集下载链接:ir.csv。
更多推荐
所有评论(0)