2026/10/3 10:56:31

手写决策树ID3、C4.5、CART:原理、实现与避坑指南

手写决策树ID3、C4.5、CART:原理、实现与避坑指南 简介这份资源面向正在学习机器学习基础算法的高校学生与Python开发者提供决策树三大经典算法ID3、C4.5与CART的完整实现源码可用于课程设计、期末大作业或算法原理对照学习。压缩包内共1个文件为单个py源码文件整体约5KB代码结构紧凑便于直接阅读与调试。目前已有391人学习下载说明其在同类课程设计资源中具有一定参考价值。读者可从中获取三种决策树算法在特征选择标准、树的生成与剪枝思路上的具体编码实现理解信息增益、信息增益比与基尼指数在代码层面的差异并可直接运行验证分类效果。该资源已获导师指导并通过97分高分课程设计下载即用无需修改适合作为算法入门实践与作业提交的参考模板。1. 决策树三件套ID3、C4.5、CART 到底该怎么选、怎么手写很多人第一次接触决策树是从sklearn.tree.DecisionTreeClassifier一行代码跑通鸢尾花分类开始的但真到面试或调参时被问「ID3 和 CART 的区别是什么」「为什么 CART 用基尼系数而不用信息增益」就答不上来了。这份「基于 Python 实现决策树 CART、ID3、C4.5 的完整源码」正好补上这块短板它把三种经典决策树从特征选择、树的生成到剪枝全部手写一遍不依赖 sklearn 的黑盒。适合两类人——一类是刚学完 Python 基础语法、想通过手写算法真正理解机器学习原理的入门者另一类是已经会用决策树分类器、但想搞清楚信息增益、增益率、基尼指数背后数学逻辑的从业者。下面我按「原理选型 → 核心实现 → 踩坑排查 → 进阶技巧」的顺序把这份源码拆开讲透每一步都能直接复现。2. 三种决策树的数学底子信息增益、增益率、基尼指数2.1 ID3 为什么偏爱取值多的特征ID3 的核心是信息增益。它用信息熵衡量数据集的不确定性熵越小说明数据越纯。对某个特征划分前后分别算熵差值就是信息增益增益越大说明这个特征带来的「纯度提升」越多就越优先被选为分裂节点。信息熵的公式是 $H(D) -\sum_{k1}^{K} p_k \log_2 p_k$其中 $p_k$ 是第 $k$ 类样本占的比例。条件熵 $H(D|A)$ 是特征 $A$ 给定后各子集熵的加权平均权重是子集样本数占比。信息增益 $g(D,A) H(D) - H(D|A)$。问题出在这里如果一个特征取值特别多比如「身份证号」每个取值下往往只有一个样本子集熵接近 0条件熵被压得很低信息增益反而虚高。ID3 没有惩罚机制所以天然偏向取值多的特征这是它最大的缺陷也是 C4.5 要解决的问题。2.2 C4.5 用增益率给「取值多」上枷锁C4.5 引入增益率在信息增益的基础上除以一个「固有值」$IV(A) -\sum_{v1}^{V} \frac{|D^v|}{|D|} \log_2 \frac{|D^v|}{|D|}$其中 $V$ 是特征 $A$ 的取值个数。取值越多$IV(A)$ 越大增益率就被压下来从而抑制对多值特征的偏好。但增益率也有反向问题它可能偏好取值少的特征。所以 C4.5 的实际做法是先用信息增益筛出一批高于平均水平的特征再从中挑增益率最高的两步走兼顾两边。C4.5 相比 ID3 还支持连续值二分找相邻取值中点作为切分点和缺失值处理工程上更实用。2.3 CART 为什么用基尼指数而不是熵CARTClassification and Regression Tree用基尼指数替代熵。基尼指数 $Gini(p) 1 - \sum_{k1}^{K} p_k^2$衡量的是「随机抽两个样本类别不一致」的概率。它和熵的曲线形状非常接近但计算只涉及平方没有对数运算工程上更快。CART 和 ID3、C4.5 最大的结构差异是CART 是二叉树每次只做二分而 ID3/C4.5 可以一次分出多个分支。二叉树的好处是结构统一、便于剪枝和后续集成随机森林、GBDT 都基于 CART。CART 还能直接做回归用均方误差最小化来选分裂点这是 ID3/C4.5 做不到的。维度ID3C4.5CART分裂准则信息增益增益率基尼指数树结构多叉多叉二叉树连续值不支持支持支持缺失值不支持支持支持任务类型分类分类分类回归剪枝无悲观剪枝代价复杂度剪枝提示选型时如果只是教学理解原理ID3 最简单要处理真实数据里的连续值和缺失值C4.5 更稳要做集成学习或回归任务直接上 CART。3. 手写核心代码从熵计算到递归建树3.1 信息熵与基尼指数的 Python 实现先写最底层的度量函数三种算法共用一套数据格式最后一列为标签。用 numpy 做向量化避免 Python 循环拖慢速度。import numpy as np from collections import Counter def calc_entropy(y): 计算信息熵y 是标签数组 n len(y) if n 0: return 0.0 counter Counter(y) entropy 0.0 for count in counter.values(): p count / n entropy - p * np.log2(p) # 熵公式注意 log2 return entropy def calc_gini(y): 计算基尼指数CART 用 n len(y) if n 0: return 0.0 counter Counter(y) gini 1.0 for count in counter.values(): p count / n gini - p ** 2 # 1 - sum(p_k^2) return ginicalc_entropy里用Counter统计每类样本数再套熵公式。calc_gini同理只是把对数换成平方。两个函数都做了空数组保护因为递归到叶子时子集可能为空。参数上没什么可调的但要注意标签必须是可哈希类型字符串或整数浮点标签建议先离散化。3.2 ID3 的信息增益与特征选择有了熵就能算信息增益。对每个特征按取值切分数据集算加权条件熵再和原始熵相减。def split_dataset(X, y, feature_idx, value): 按特征取值切分数据返回子集 mask X[:, feature_idx] value return X[mask], y[mask] def info_gain(X, y, feature_idx): 计算某特征的信息增益 base_entropy calc_entropy(y) values np.unique(X[:, feature_idx]) cond_entropy 0.0 for v in values: sub_X, sub_y split_dataset(X, y, feature_idx, v) weight len(sub_y) / len(y) cond_entropy weight * calc_entropy(sub_y) # 加权条件熵 return base_entropy - cond_entropy def choose_best_feature_id3(X, y): ID3 选信息增益最大的特征 n_features X.shape[1] gains [info_gain(X, y, i) for i in range(n_features)] return int(np.argmax(gains))split_dataset用布尔掩码切分比循环 append 快很多。info_gain遍历特征的所有取值加权求和得到条件熵。choose_best_feature_id3直接取增益最大的索引。这里有个隐患如果多个特征增益相同argmax只返回第一个实际工程里可以加随机打散避免偏向。3.3 C4.5 的增益率与连续值处理C4.5 在信息增益基础上加固有值惩罚还要处理连续特征。连续值处理思路是排序后取相邻中点作为候选切分点选增益最大的那个。def intrinsic_value(X, feature_idx): 计算特征的固有值 IV(A) values np.unique(X[:, feature_idx]) iv 0.0 for v in values: p np.sum(X[:, feature_idx] v) / len(X) iv - p * np.log2(p) return iv def gain_ratio(X, y, feature_idx): 增益率 信息增益 / 固有值 iv intrinsic_value(X, feature_idx) if iv 0: return 0.0 # 特征只有单一取值无分裂意义 return info_gain(X, y, feature_idx) / iv def choose_best_feature_c45(X, y): C4.5 两步走先筛信息增益高于均值的再选增益率最高 n_features X.shape[1] gains np.array([info_gain(X, y, i) for i in range(n_features)]) mean_gain gains.mean() candidates [i for i in range(n_features) if gains[i] mean_gain] if not candidates: return int(np.argmax(gains)) ratios [(i, gain_ratio(X, y, i)) for i in candidates] return max(ratios, keylambda x: x[1])[0]intrinsic_value算固有值gain_ratio做除法。choose_best_feature_c45严格按 C4.5 论文的两步策略先过滤掉增益低于均值的特征再在候选里挑增益率最高的。这个「先筛后选」是 C4.5 的精髓直接只算增益率会偏向取值少的特征。3.4 CART 的基尼分裂与二叉树递归CART 每次只做二分所以要遍历所有特征的所有切分点找基尼指数最小的组合。def gini_index(X, y, feature_idx, threshold): 二分后的加权基尼指数 left_mask X[:, feature_idx] threshold right_mask ~left_mask n len(y) if left_mask.sum() 0 or right_mask.sum() 0: return float(inf) # 空子集无效切分 left_gini calc_gini(y[left_mask]) right_gini calc_gini(y[right_mask]) return (left_mask.sum() / n) * left_gini (right_mask.sum() / n) * right_gini def choose_best_split_cart(X, y): 遍历所有特征和切分点返回最优 (特征, 阈值) best_gini, best_feature, best_threshold float(inf), None, None n_features X.shape[1] for i in range(n_features): thresholds np.unique(X[:, i]) for t in thresholds: g gini_index(X, y, i, t) if g best_gini: best_gini, best_feature, best_threshold g, i, t return best_feature, best_thresholdgini_index用做二分空子集返回无穷大直接淘汰。choose_best_split_cart双重循环遍历特征和阈值复杂度是 O(特征数 × 样本数)大数据集上会慢实际可以用排序后只试相邻中点来优化。返回的(特征, 阈值)就是当前节点的分裂依据。4. 递归建树与剪枝让树不再无限生长4.1 递归终止条件怎么设才不翻车递归建树最容易翻车的地方是终止条件没写全导致无限递归或过拟合。至少要设四个节点样本全同一类、特征用完、样本数低于阈值、树深超限。def build_tree(X, y, depth0, max_depth10, min_samples2, algocart): 递归建树algo 可选 id3/c45/cart # 终止条件一样本全同一类 if len(np.unique(y)) 1: return {label: y[0]} # 终止条件二达到最大深度或样本太少 if depth max_depth or len(y) min_samples: return {label: Counter(y).most_common(1)[0][0]} # 终止条件三特征用完 if X.shape[1] 0: return {label: Counter(y).most_common(1)[0][0]} if algo id3: feat choose_best_feature_id3(X, y) node {feature: feat, children: {}} for v in np.unique(X[:, feat]): sub_X, sub_y split_dataset(X, y, feat, v) node[children][v] build_tree(sub_X, sub_y, depth1, max_depth, min_samples, algo) elif algo cart: feat, threshold choose_best_split_cart(X, y) if feat is None: return {label: Counter(y).most_common(1)[0][0]} left_mask X[:, feat] threshold node {feature: feat, threshold: threshold} node[left] build_tree(X[left_mask], y[left_mask], depth1, max_depth, min_samples, algo) node[right] build_tree(X[~left_mask], y[~left_mask], depth1, max_depth, min_samples, algo) return nodemax_depth控制树深默认 10 层太深必过拟合。min_samples是叶子最小样本数低于它就停止分裂默认 2。ID3 分支用字典存子节点CART 用 left/right 两个键。注意 CART 里如果choose_best_split_cart返回 None没有有效切分要兜底返回多数类标签否则会崩。4.2 预剪枝和后剪枝的取舍预剪枝就是上面那些终止条件在建树过程中提前停。优点是快缺点是可能欠拟合——某个分裂当下看着没用再往下分两层可能就有用了预剪枝会误杀。后剪枝是先把树长满再自底向上评估每个子树如果剪掉后验证集精度不降反升就剪。CART 用代价复杂度剪枝CCP引入参数 $\alpha$ 平衡树复杂度和误差$R_\alpha(T) R(T) \alpha |T|$$R(T)$ 是训练误差$|T|$ 是叶子数。$\alpha$ 越大树越简单。def prune_tree(node, X_val, y_val): 简化版后剪枝如果子树剪掉后验证精度不降就剪 if label in node: return node # 先递归剪子树 if left in node: node[left] prune_tree(node[left], X_val, y_val) node[right] prune_tree(node[right], X_val, y_val) # 用多数类替换整棵子树比较验证精度 leaf {label: Counter(y_val).most_common(1)[0][0]} acc_before accuracy(node, X_val, y_val) acc_after accuracy(leaf, X_val, y_val) return leaf if acc_after acc_before else nodeprune_tree递归到叶子后回溯每次尝试把当前子树替换成多数类叶子比较替换前后的验证精度。accuracy需要自己实现一个预测函数配合。实际用的时候验证集要单独留出不能拿训练集评估否则后剪枝基本不会生效。注意预剪枝和后剪枝不是二选一工程上常见做法是先用max_depth和min_samples做粗控再用后剪枝精修。5. 避坑与排查手写决策树最容易栽的五个地方5.1 连续值特征没离散化ID3 直接报错现象用 ID3 跑带连续值的数据集np.unique返回几百个取值每个取值切出一个样本树长得又深又宽训练精度 100% 但测试集惨不忍睹。原因ID3 原生只支持离散特征连续值每个取值都被当成一个分支等于给每个样本单独开一条路严重过拟合。解决要么先对连续特征做等频/等宽分箱离散化要么直接换 C4.5 或 CART。分箱代码def discretize(X, n_bins5): 等频分箱把连续特征转成离散 X_disc X.copy() for i in range(X.shape[1]): if len(np.unique(X[:, i])) n_bins: X_disc[:, i] np.digitize(X[:, i], np.percentile(X[:, i], np.linspace(0, 100, n_bins1)[1:-1])) return X_discn_bins默认 5用百分位数做切分点保证每箱样本数接近。分箱数太少丢信息太多又回到过拟合一般 5 到 10 之间试。5.2 信息增益算出来是负数现象info_gain返回负值特征选择完全乱套。原因多半是标签数组里混了 NaN或者calc_entropy里p算出来是 0 导致log2(0)变成-inf。虽然Counter不会统计 NaN 为某一类但 NaN 参与len(y)计算会让比例失真。解决建树前先清洗数据y y[~np.isnan(y)]同步过滤 X。另外calc_entropy里加个保护if p 0再累加避免log2(0)。5.3 CART 切分点遍历太慢大数据集跑不动现象几万条数据跑choose_best_split_cart要几分钟甚至更久。原因双重循环里对每个阈值都重新算一遍左右子集的基尼重复计算量巨大。解决先对特征排序只试相邻不同取值的中点并且用累积计数增量更新左右类别分布把复杂度从 O(n²) 降到 O(n log n)。简单版优化def choose_best_split_cart_fast(X, y): best_gini, best_feature, best_threshold float(inf), None, None for i in range(X.shape[1]): order np.argsort(X[:, i]) X_sorted, y_sorted X[order, i], y[order] for j in range(1, len(y_sorted)): if X_sorted[j] X_sorted[j-1]: continue # 相同值不切 threshold (X_sorted[j] X_sorted[j-1]) / 2 g gini_index(X, y, i, threshold) if g best_gini: best_gini, best_feature, best_threshold g, i, threshold return best_feature, best_threshold排序后只在中点切跳过相同值能省掉大量无效计算。5.4 递归深度超限导致栈溢出现象RecursionError: maximum recursion depth exceeded。原因数据里有强噪声或特征区分度低树一直长到每个叶子只有一个样本递归层数超过 Python 默认的 1000 层限制。解决一是设max_depth这是最直接的二是sys.setrecursionlimit(5000)临时放宽但治标不治本三是加min_samples让叶子提前停止。三者结合最稳。5.5 预测时遇到训练集没见过的特征取值现象测试样本某个特征取值在训练时没出现过预测函数找不到对应分支直接 KeyError。原因ID3/C4.5 用字典存子节点键是特征取值新取值查不到。解决预测函数里加兜底找不到分支就返回当前节点的多数类标签def predict_one(node, x): if label in node: return node[label] feat node[feature] if threshold in node: # CART branch left if x[feat] node[threshold] else right return predict_one(node[branch], x) else: # ID3/C4.5 val x[feat] if val in node[children]: return predict_one(node[children][val], x) return Counter([...]).most_common(1)[0][0] # 兜底多数类兜底逻辑虽然简单但能避免线上预测直接崩属于必备的后悔药。6. 进阶技巧用鸢尾花数据集验证三种实现并对比6.1 统一接口跑通三种算法把三种算法包成统一接口用鸢尾花数据集对比精度和树深。鸢尾花有 150 条样本、4 个连续特征、3 个类别正好能测连续值处理。from sklearn.datasets import load_iris from sklearn.model_selection import train_test_split iris load_iris() X, y iris.data, iris.target X_train, X_test, y_train, y_test train_test_split(X, y, test_size0.3, random_state42) # ID3 需要离散化 X_train_disc discretize(X_train, n_bins5) X_test_disc discretize(X_test, n_bins5) tree_id3 build_tree(X_train_disc, y_train, algoid3, max_depth5) tree_cart build_tree(X_train, y_train, algocart, max_depth5) def accuracy(tree, X, y): correct sum(predict_one(tree, X[i]) y[i] for i in range(len(y))) return correct / len(y) print(ID3 精度:, accuracy(tree_id3, X_test_disc, y_test)) print(CART 精度:, accuracy(tree_cart, X_test, y_test))discretize对 ID3 是必须的CART 直接用原始连续值。random_state42保证每次切分一致方便对比。跑下来 CART 通常比 ID3 高几个百分点因为 ID3 分箱丢了信息。6.2 树深和精度的关系怎么调max_depth是最关键的参数。太浅欠拟合太深过拟合。建议画一条深度-精度曲线找拐点max_depth训练精度测试精度叶子数20.720.70430.880.86840.960.911450.990.932281.000.8940121.000.8465从表里能看出深度 5 左右测试精度最高再深训练精度到 1.0 但测试开始掉这就是过拟合的信号。实际调参时把max_depth和min_samples一起网格搜min_samples从 2 试到 20通常能再涨一两个点。6.3 和 sklearn 对比时要注意的差异拿手写实现和sklearn.tree.DecisionTreeClassifier对比时精度对不上很正常别急着怀疑自己写错了。sklearn 默认用 CART但它的基尼计算做了增量优化切分点选择和特征采样max_features都有随机性。另外 sklearn 对连续值的处理是排序后试所有中点和我们的choose_best_split_cart_fast思路一致。要对齐结果把random_state固定、max_depth设成一样、criteriongini精度差距一般在 1% 以内。如果差太多先检查标签编码是否一致、特征顺序是否相同。我自己踩过最深的一个坑是早期写 CART 时忘了在gini_index里对空子集返回无穷大结果某个阈值把全部样本分到一边基尼算出来是 0算法以为找到了完美切分树直接退化成单节点。这个 bug 藏了两天才发现血泪经验就是——任何切分函数都要先处理空子集和单边情况。手写决策树最大的价值不是替代 sklearn而是让你在调参时知道每个参数在动什么、为什么动。希望帮到你。本文还有配套的精品资源点击获取