1. 决策树为什么值得从原理重新啃一遍
做机器学习这几年,我发现一个很有意思的现象:很多人用随机森林、XGBoost用得飞起,但被问到“一棵树到底是怎么长出来的”,往往只能说出个大概。决策树和随机森林的区别、决策树信息增益到底怎么算、为什么CART能同时做分类和回归,这些细节才是真正拉开水平的地方。今天这篇就当是一次复盘,把ID3、C4.5、CART这三大算法从数学原理到代码实现完整过一遍,顺便把我在实际项目中踩过的坑也一起说了。
这篇内容适合刚入门决策树算法的新手,也适合那些已经把sklearn用得滚瓜烂熟、但想回头把底子补扎实的同学。我尽量用大白话讲清楚每一个公式背后的直觉,再用手算例子让你能跟着推一遍。读完你至少会明白三件事:一棵树是怎么选择分裂特征的、三种算法的本质区别到底在哪、以及实战中那些参数到底在控制什么。
2. 一棵树是怎么“长”出来的
2.1 决策树在做什么
决策树这个名字听着唬人,其实它的本质就是一组“if-else”规则的嵌套。想象你在判断“今天要不要出门打球”,你会先看天气,如果下雨就不去,如果晴天再看温度,温度太高也不去,温度合适再看风力……每一次判断都在把候选答案分成两支或多支,最后落到一个具体结论上。计算机里的决策树干的就是这件事,只不过它面对的不是天气、温度这种直观变量,而是表格里的特征列。
问题来了:现实数据里可能有几十个特征,哪些特征适合放在树的上层先判断?这就需要一个衡量标准。ID3用的是信息增益,C4.5用信息增益率,CART用Gini指数。这三个标准本质上都在回答同一个问题:用某个特征做划分之后,数据的“混乱程度”降低了多少。降低得越多,说明这个特征越能区分不同的类别,就越应该优先拿来分裂。
2.2 从一堆特征里怎么选“先问谁”
举个具体例子:假设有一个关于用户是否购买商品的表格,特征是年龄、收入、是否学生,标签是“买/不买”。如果按“是否学生”来划分,买和不买的人可能被分得很干净;如果按“年龄”来划分,可能分完之后每个子集里还是有买有不买,跟没分差不多。前者明显是更好的分裂方式,因为它让每个子集内部的纯度更高了。
纯度怎么量化?“信息熵”就是干这个的。熵这个概念来自信息论,它衡量的是一个集合内部的不确定性。如果一堆样本全是同一个类别,熵就是0,表示完全确定;如果两类各占一半,熵就是1,表示最混乱。决策树每一次分裂,都在想办法让分裂后的子集熵值总和尽可能低,也就是让每个子集尽可能“纯”。
这里有个容易搞混的点:信息增益算的是“分裂前熵减去分裂后加权熵”的差值,差得越多说明这个特征带来的纯度提升越大。但增益大的特征不一定最好,因为类别多的特征天然占便宜,C4.5的信息增益率就是专门来治这个毛病的,后面我会细说。
3. ID3算法:信息增益是怎么算出来的
3.1 信息熵与条件熵的手算过程
ID3是决策树最早期的代表算法,由Quinlan在1986年提出。它的分裂准则是信息增益,计算分三步。
第一步,算分裂前的熵。假设我们有14个样本,9个买、5个不买,熵的计算公式是:
import math def entropy(positive, negative): total = positive + negative p_pos = positive / total p_neg = negative / total if p_pos == 0 or p_neg == 0: return 0 return -(p_pos * math.log2(p_pos) + p_neg * math.log2(p_neg)) print(entropy(9, 5))算出来约等于0.940,这就是根节点的初始熵。
第二步,按某个特征划分后,算条件熵。比如特征是“天气”,有三个取值:晴天5个(其中买2个)、阴天4个(其中买4个)、雨天5个(其中买3个)。条件熵就是把每个子集的熵按样本占比加权求和:
# 晴天子集:2买3不买 e_sunny = entropy(2, 3) # 阴天子集:4买0不买 e_overcast = entropy(4, 0) # 雨天子集:3买2不买 e_rainy = entropy(3, 2) cond_entropy = 5/14 * e_sunny + 4/14 * e_overcast + 5/14 * e_rainy print(cond_entropy)第三步,信息增益 = 0.940 - 条件熵。哪个特征算出来的增益最大,就选它作为当前节点的分裂特征,然后对每个分支重复这个过程,直到所有样本都属于同一类别,或者没有特征可用。
3.2 ID3的致命伤:偏爱取值多的特征
ID3在实际应用中很快暴露出一个问题:它天然偏爱取值数量多的特征。比如把“身份证号”作为一个特征,每个人都有一个唯一取值,按它划分后每个子集只有一个样本,条件熵直接降到0,信息增益拉满,决策树一定会优先选它。但这个特征根本没有泛化能力,拿到新样本上毫无意义。
这就是过拟合的雏形。树为了把训练数据学得完美,长得又深又碎,一到测试集上就露馅。我当时第一次手写ID3的时候,真就遇到过这种情况,树长了好几层,训练集准确率接近100%,测试集一测直接掉到70%。后来查资料才明白,这不是代码写错了,是算法本身的倾向性导致的。
3.3 ID3的其他局限
除了偏爱多取值特征,ID3只能处理离散型特征。连续特征比如年龄、收入,如果不事先做离散化,ID3根本没法用。另外它不能处理缺失值,样本里有空值就得先扔掉,这在真实业务数据里非常麻烦。
还有一个问题:ID3不支持剪枝。树会长到“完美拟合训练集”才会停下来,这在数据有噪声时特别致命。比如有两个样本特征一模一样但标签不同,ID3会为了区分它们硬生生长出一些没有统计意义的节点。所以后来的C4.5几乎是在给ID3打补丁:加了信息增益率、支持连续特征、支持缺失值、还引入了剪枝机制。
4. C4.5算法:补丁安完的进阶版
4.1 信息增益率怎么解决“偏爱多取值”的问题
C4.5是Quinlan在1993年提出的,算是ID3的全面升级版。它不再直接用信息增益,而是用信息增益率——增益除以一个“惩罚项”,这个惩罚项是特征本身的熵,也就是按该特征取值划分后,取值分布的混乱程度。
举个例子,按“身份证号”划分时,因为每个取值只对应一个样本,特征的熵会非常高,增益率就被压下去了。而像“天气”这种只有三个取值的特征,它的特征熵不高,增益率就会体现得更合理。这样一来,C4.5就不会盲目选择取值多的特征了。
增益率的计算公式是:
def gain_ratio(information_gain, feature_entropy): if feature_entropy == 0: return 0 return information_gain / feature_entropy这里有个细节容易踩坑:如果某个特征只有一个取值,特征熵是0,增益率会变成无穷大。所以实际代码里一般会对取值数量做限制,或者直接跳过这种特征。
4.2 连续特征和缺失值终于能用了
连续特征的处理思路是二分法。比如特征“年龄”有一堆取值,排序后取相邻两个值的中间点作为候选切分点,然后计算每个切分点的信息增益率,选最大的那个作为分裂点。这意味着连续特征在C4.5中只会产生二叉分裂,而离散特征仍然可以多叉分裂。
缺失值的处理则分两种情况:一种是计算信息增益率时,只统计没有缺失的样本;另一种是样本进入某个分支时,如果特征缺失,就把它同时分到所有子节点,并加一个权重系数。这个机制后来在XGBoost里也有类似的实现思路,所以理解C4.5对理解后续算法很有帮助。
4.3 剪枝策略:让树别长那么野
C4.5引入了两种剪枝:预剪枝和后剪枝。预剪枝是在树生长过程中提前停止,比如信息增益率小于某个阈值就不再分裂;后剪枝是先让树长完整,再从下往上把某些子树替换成叶子节点,用验证集来判断替换后准确率会不会下降。
我在实际项目中试下来,后剪枝的效果通常比预剪枝好,因为它有全局视角,知道哪些分支是“局部看着有用、整体没啥用”的。预剪枝的参数很难调,设太早树就变成傻大个,设太晚又等于没剪。所以如果你用C4.5这类算法,建议优先考虑后剪枝。
5. CART算法:二叉树的极致
5.1 Gini指数比熵好在哪
CART(Classification and Regression Tree)是Breiman在1984年提出的,也是现在scikit-learn里DecisionTreeClassifier的默认实现。它的核心区别有两个:一是始终做二叉分裂,二是不用信息熵而用Gini指数。
Gini指数的计算比熵简单,没有对数运算。它衡量的是“从集合里随机抽两个样本,它们类别不同的概率”。公式是:
def gini(labels): total = len(labels) if total == 0: return 0 prob = {} for label in labels: prob[label] = prob.get(label, 0) + 1 impurity = 1 for count in prob.values(): impurity -= (count / total) ** 2 return impurityGini指数越小,集合越纯。按特征分裂时,CART计算的是分裂前后Gini指数的加权差,选差值最大的特征和切分点。
有人会问,熵和Gini到底哪个好?实践中两者效果非常接近,但Gini计算更快,因为不需要算对数。当类别很多时效率差异会明显一些。所以sklearn默认用Gini不是没有道理的。
5.2 回归树又是怎么回事
CART不仅能做分类,还能做回归。区别在于分裂标准从Gini指数换成了均方误差(MSE)。回归树每个叶子节点的输出不再是类别,而是该叶子下所有样本目标值的均值。
分裂时,算法会尝试所有特征的候选切分点,计算切分后左右两个子集的MSE之和,选择让整体MSE最小的那个切分点。比如预测房价,特征是面积,按面积中位数切一刀后,如果左右两边房子价格分别都差不多,MSE就小,说明这个切分效果好。
回归树的一个明显弱点是它输出的是阶梯状的预测值,不够平滑。同一个叶子里的样本预测值完全相同,所以树深不够时会看到一片一片的平台。实践中可以用随机森林或梯度提升树来平抑这个问题。
5.3 CART的剪枝:代价复杂度剪枝
CART的后剪枝用的是代价复杂度剪枝(Cost Complexity Pruning)。思路是定义一个损失函数:R(T) + α|T|,其中R(T)是预测误差,|T|是叶子节点数量,α是惩罚系数。树越深误差越小但叶子越多,α越大就越偏向小树。sklearn里的ccp_alpha参数就是干这个的。
代价复杂度剪枝的精髓在于它会生成一串不同α对应的子树,然后用交叉验证选最优的那棵。这是一种很优雅的做法,因为它把树的复杂度当成一个可调的超参数来处理了。我调ccp_alpha的时候喜欢画一条曲线,横轴是α,纵轴是验证集准确率,找一个曲线开始明显下滑之前的位置,就是合适的α。
6. 三大算法对比:一张表看懂
表格是最直观的对比方式,我把ID3、C4.5、CART的核心差异列出来:
| 对比维度 | ID3 | C4.5 | CART |
|---|---|---|---|
| 分裂准则 | 信息增益 | 信息增益率 | Gini指数 / MSE |
| 树结构 | 多叉树 | 多叉树 | 二叉树 |
| 连续特征 | 不支持 | 支持(二分切分) | 支持(二分切分) |
| 缺失值处理 | 不支持 | 支持 | 支持 |
| 剪枝 | 不支持 | 预剪枝+后剪枝 | 代价复杂度剪枝 |
| 任务类型 | 分类 | 分类 | 分类+回归 |
| 常见库实现 | 少见 | 少见 | sklearn、Spark MLlib |
为什么现在的工业界几乎都默认用CART?原因很简单:二叉树在实现上更简洁,每个节点的分裂逻辑只有“左还是右”,而多叉树可以直接用二叉形式表示且分类效果基本相同。再加上CART天然支持回归,一个算法通吃两个任务,工程上太方便了。
从这里也能看出算法演进的逻辑线:ID3提出了“用信息增益选特征”的基本框架,C4.5修补了它的各种缺陷,CART则把树结构简化到了极致并扩展了应用范围。三者之间是迭代关系,不是互斥关系。
7. 从零手写一个简化版决策树
7.1 代码实现:核心结构
说了这么多理论,接下来上代码。下面是一个简化版的CART分类树实现,只保留核心分裂和建树逻辑,方便你理解树的生长过程:
import numpy as np from collections import Counter class Node: def __init__(self, feature_idx=None, threshold=None, left=None, right=None, value=None): self.feature_idx = feature_idx self.threshold = threshold self.left = left self.right = right self.value = value # 叶子节点的预测类别 def gini(labels): total = len(labels) if total == 0: return 0 prob = Counter(labels) impurity = 1 for count in prob.values(): impurity -= (count / total) ** 2 return impurity def split_dataset(X, y, feature_idx, threshold): left_mask = X[:, feature_idx] <= threshold right_mask = ~left_mask return X[left_mask], y[left_mask], X[right_mask], y[right_mask] def best_split(X, y): best_gain = -1 best_feature = None best_threshold = None n_features = X.shape[1] parent_gini = gini(y) for f_idx in range(n_features): values = np.unique(X[:, f_idx]) for i in range(len(values) - 1): threshold = (values[i] + values[i + 1]) / 2 _, y_left, _, y_right = split_dataset(X, y, f_idx, threshold) if len(y_left) == 0 or len(y_right) == 0: continue weighted_gini = (len(y_left) * gini(y_left) + len(y_right) * gini(y_right)) / len(y) gain = parent_gini - weighted_gini if gain > best_gain: best_gain = gain best_feature = f_idx best_threshold = threshold return best_feature, best_threshold def build_tree(X, y, max_depth=3, depth=0): if len(np.unique(y)) == 1 or depth >= max_depth or X.shape[0] == 0: return Node(value=Counter(y).most_common(1)[0][0]) feature_idx, threshold = best_split(X, y) if feature_idx is None: return Node(value=Counter(y).most_common(1)[0][0]) X_left, y_left, X_right, y_right = split_dataset(X, y, feature_idx, threshold) left_node = build_tree(X_left, y_left, max_depth, depth + 1) right_node = build_tree(X_right, y_right, max_depth, depth + 1) return Node(feature_idx=feature_idx, threshold=threshold, left=left_node, right=right_node) def predict_one(node, x): if node.value is not None: return node.value if x[node.feature_idx] <= node.threshold: return predict_one(node.left, x) else: return predict_one(node.right, x) def predict(tree, X): return np.array([predict_one(tree, x) for x in X])这段代码不到60行,但完整跑通了“找最佳分裂点 → 递归建树 → 预测”的整个流程。你可以用一个简单的数据集直接测试:
from sklearn.datasets import load_iris from sklearn.model_selection import train_test_split from sklearn.metrics import accuracy_score data = load_iris() X_train, X_test, y_train, y_test = train_test_split(data.data, data.target, test_size=0.2, random_state=42) tree = build_tree(X_train, y_train, max_depth=3) y_pred = predict(tree, X_test) print("准确率:", accuracy_score(y_test, y_pred))用鸢尾花数据集跑一遍,准确率大约在93%左右,和sklearn默认参数下的结果差距很小。这说明核心逻辑是对的。
7.2 关键点解读:为什么切分阈值这么取
代码里取候选切分点时用了相邻两个特征值的平均值,这就是CART处理连续特征的经典做法。为什么取中点而不是直接用特征值本身?因为特征值可能是不连续的,取中点能保证两个分支都有样本覆盖,不会出现某个分支为空的情况。
还有一个细节:split_dataset里用了<=和>,这是CART二叉分裂的标准写法。如果你实现的是多叉树,就不需要threshold这个东西,直接按类别取值分多个分支就行。
树的停止条件也很重要。这段代码里只用了两个:类别纯净、达到最大深度。实际产品里还需要加“最小样本数”“最小不纯度下降量”等条件,否则树容易过度生长。
7.3 手写版和sklearn版的差距在哪
手写代码的意义在于理解算法,但真到工业应用,sklearn的版本考虑了太多工程细节:特征预排序、并行计算、样本权重、类别权重、代价复杂度剪枝路径,等等。它的效率远高于手写版,而且经过了海量用户验证,稳定性有保障。
我建议你把手写版当作“教学模型”来用,理解每一行在干什么就好。实际项目里还是老老实实调sklearn,不要重复造轮子,但造过轮子的人调起参来,思路会清晰很多。
8. 实战调参与常见问题
8.1 过拟合和欠拟合怎么平衡
决策树最大的优点是好解释、好可视化,最大的缺点就是容易过拟合。一个常见现象是训练集准确率99%、测试集75%,这基本就是树太深了。解决办法无非几个方向:限制max_depth、增大min_samples_split、增大min_samples_leaf、调ccp_alpha。
我个人的调参习惯是先用默认参数跑一遍,看训练集和测试集的差距。差距大就优先限制深度和叶节点最小样本数,差距小但整体准确率低,那就是欠拟合了,应该加大max_depth或者换更强的模型(比如随机森林、决策树回归里的梯度提升树)。
有一点值得注意:决策树和随机森林是两回事,随机森林是很多棵树的集成,用bagging来降低方差。如果你发现单棵决策树过拟合严重,可以先试调参;如果调了半天还是不行,果断换随机森林,通常会有质的提升。
8.2 特征重要性怎么看
决策树有一个天然副产品:特征重要性。每次分裂时,某个特征带来的不纯度下降量会被累计,最后归一化就是特征重要性的分数。在sklearn里直接用model.feature_importances_就能拿出来。
但这个分数有个陷阱:它偏向高基数特征,也就是取值多、切分机会多的特征。这跟ID3偏爱多取值特征的问题是同一根源。所以看特征重要性排序时,不要完全相信单棵树的结果,用随机森林或多次交叉验证后的平均结果会更可靠。
8.3 连续值和缺失值实战中的处理
虽然CART理论上能处理连续值,但特征之间的量纲差异会影响阈值选择。比如一个特征是“收入”单位是元,另一个特征是“年龄”单位是岁,收入特征搜索空间大得多,可能获得更多切分机会。所以实战中最好先做标准化或归一化,虽然树模型对单调变换不敏感,但对取值范围差异还是有反应的。
缺失值方面,sklearn里的DecisionTreeClassifier默认不处理缺失值,需要你自己填充。简单做法是用中位数或众数填充,更讲究一点可以用模型预测缺失值,或者用带缺失值支持的工具库。实际业务里缺失值多的话,树模型的效果会很受影响,这个坑我踩过不止一次。
8.4 决策树回归的应用场景
很多人以为决策树只能做分类,其实决策树回归在很多场景下都很好用。比如预测用户活跃度、预估订单时长、评估设备寿命,这些连续值的预测任务都可以用CART回归树来建模。
决策树回归和线性回归的差别在于,决策树回归不假设数据存在线性关系,可以捕捉非线性模式。代价就是预测结果不平滑,而且外推能力差——如果测试样本的特征值超出了训练集的范围,树只能输出训练集里最后那个叶子的均值,做不到线性回归那种“沿着趋势往外推”的效果。所以在做时间序列预测或者需要外推的场景,要谨慎使用决策树回归。