尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

决策树三大经典算法ID3、C4.5、CART的Python实现与原理详解

决策树三大经典算法ID3、C4.5、CART的Python实现与原理详解 简介压缩包提供了一份面向Python初学者的决策树算法实现围绕ID3、C4.5与CART三种经典算法展开演示如何基于信息熵、信息增益比和基尼不纯度进行特征划分、建树与预测。代码采用模块化结构以鸢尾花数据集为实例完整覆盖数据导入、模型构建、训练与预测流程包体共9个文件包括6个Python源码、2个pyc编译文件及1份数据集压缩包仅14KB轻量易读适合初学者快速上手。目前已有389人学习下载。运行时可直观对比三种算法的分裂标准ID3倾向多叉树且难以处理连续值C4.5引入增益比并支持连续属性和缺失值CART使用基尼系数生成二叉树还能用于回归。通过动手调试和修改参数能加深对决策树过拟合、剪枝与特征选择的理解为后续学习scikit-learn及实际项目选型打下基础。1. 决策树三种经典算法一篇能跑通ID3、C4.5、CART的Python实现笔记很多人用sklearn调用一个决策树分类器只需要三行代码但一旦面试官问“信息增益怎么算”、课程设计要求你手写算法才发现自己只会fit和predict。标题里的“决策树三种经典算法实现.rar”这类资源包通常装的是ID3、C4.5、CART三个算法的Python源码加一份数据但直接拿源码容易被里面绕来绕去的边界处理搞晕。这篇文章不绕弯按我自己在课程设计和竞赛里反复验证过的一套方案把三种算法的原理、特征怎么选、代码怎么写、坑在哪里全部拆开。适合正在学机器学习基础、需要交决策树作业或者想彻底看懂sklearn决策树背后逻辑的人。全程用Python先讲清楚每个分裂准则再给完整实现最后教你调参和避坑。2. 先分清三种算法信息增益、增益率、基尼指数到底差在哪2.1 ID3信息增益如何选择最优特征ID3是最早也最直觉的决策树算法。它的核心是“信息增益”每次分裂都找一个特征让分裂后各个子节点的数据尽可能“纯”。判断纯度用信息熵Entropy。对数据集D假设有K个类别每一类占比为p_k熵的计算公式是H(D) -Σ p_k * log2(p_k)熵越大表示混乱程度越高。加入特征A后按A的取值把D切成若干子集D_i每个子集熵的加权平均就是条件熵H(D|A)。两者相减就是特征A带来的信息增益gain(D, A) H(D) - H(D|A)ID3遍历所有特征选增益最大的那个。写成Python就是下面这个函数这也是后面所有实现的基础。import numpy as np def entropy(labels): # labels: 一维数组当前节点的全部类别标签 classes, counts np.unique(labels, return_countsTrue) probs counts / len(labels) # 概率为0的类别不需要参与否则log2(0)会变成-inf return -np.sum(probs * np.log2(probs))这个函数的逻辑很直白用np.unique统计每个类别的数量除以总数得到概率然后按公式累加。注意p为0的时候要跳过否则代码会报“divide by zero”的RuntimeWarning最后得到一个inf。这个细节很多人第一次写都会踩到。计算增益时还要对每个特征单独算条件熵。def info_gain(data, feature, target): # data: pandas DataFrame, target: 标签列名 base_entropy entropy(data[target]) values data[feature].unique() new_entropy 0.0 for v in values: subset data[data[feature] v] prob len(subset) / len(data) new_entropy prob * entropy(subset[target]) return base_entropy - new_entropy这里的参数就两个feature是当前考察的特征列target是标签列。如果某个特征取值特别多比如“编号”每个子集只有一个样本条件熵变成0信息增益就等于基础熵ID3会无脑选它。这就是ID3的先天缺陷也是C4.5要解决的问题。2.2 C4.5信息增益率修正了ID3的什么毛病C4.5在ID3上做了三件重要的事改用信息增益率支持连续特征处理缺失值。信息增益率不是直接用增益而是先计算分裂信息Split Informationsplit_info(D, A) -Σ |D_i|/|D| * log2(|D_i|/|D|)然后增益率等于增益除以分裂信息gain_ratio gain / split_info分裂信息衡量的是特征本身的取值丰富程度。取值越多split_info越大会把过高的信息增益压下去。这个修正让算法不再偏爱取值多的特征。连续特征处理是C4.5的另一个关键。比如“温度”是数值不能直接按值分叉。常见做法是先把特征值排序取相邻两个值的中点作为候选阈值然后像离散特征一样计算每个阈值下的信息增益率。假设有n个样本最多产生n-1个候选阈值。数据量大的时候可以只取分位点加速。def gain_ratio_for_threshold(data, feature, threshold, target): # 按阈值把连续特征切成左右两半 left data[data[feature] threshold] right data[data[feature] threshold] if len(left) 0 or len(right) 0: return -np.inf base_entropy entropy(data[target]) new_entropy (len(left) / len(data) * entropy(left[target]) len(right) / len(data) * entropy(right[target])) gain base_entropy - new_entropy split_info - (len(left) / len(data) * np.log2(len(left) / len(data)) len(right) / len(data) * np.log2(len(right) / len(data))) if split_info 0: return -np.inf return gain / split_info这个函数展示了C4.5对连续特征的核心操作每次只把数据分成“小于等于阈值”和“大于阈值”两部分。注意split_info变成0的情况说明阈值把数据都切在一边这种情况要直接返回负无穷否则除零报错。2.3 CART基尼指数为什么更适合回归和剪枝CARTClassification And Regression Tree用的不是熵是基尼指数Gini Index。基尼指数的意义是从集合里随机抽两个样本它们的类别不一致的概率。公式Gini(D) 1 - Σ p_k^2基尼指数越小集合越纯。对特征A基尼指数按子集加权Gini_idx(D, A) Σ |D_i|/|D| * Gini(D_i)CART和ID3/C4.5最大的不同是它永远只做二叉树。不管特征是离散还是连续每次分裂只把数据切成两堆。离散特征的做法是考虑“取值为v”和“取值不为v”两个子集连续特征同样是找阈值二分。这种二叉结构方便递归也方便做后剪枝还能直接扩展到回归树把基尼指数换成均方误差。def gini(labels): classes, counts np.unique(labels, return_countsTrue) probs counts / len(labels) return 1 - np.sum(probs ** 2)和熵相比基尼指数不需要算log计算速度更快。这也是sklearn默认用CART的原因之一。2.4 选型对照表什么时候用哪种算法分裂准则树结构连续特征缺失值典型场景ID3信息增益多叉不支持不支持教学、理解熵的概念C4.5信息增益率多叉支持支持小规模数据、课程设计CART基尼指数二叉支持支持需额外实现sklearn默认、工业应用如果你只是交作业ID3最小可用如果数据里有连续数值直接上C4.5或CART如果想和sklearn的DecisionTreeClassifier对齐优先CART。但我的建议是三个都实现因为它们的框架完全一样只是“选哪个特征分裂”这个函数不同。理解了这一点后面代码就能少写一半。3. 准备Python环境与数据手写决策树前先解决这四件套3.1 用纯Python还是sklearn我的选择是手写核心标题里的资源包既然叫“决策树三种经典算法实现”重点在“实现”这两个字。有的包里写的其实只是调了sklearn的DecisionTreeClassifier这就失去了意义。我的做法是核心分裂和建树逻辑全部手写只用NumPy做数组运算、Pandas做数据读取sklearn只用来加载数据集和最后算准确率。这样既不是纯调库又能借助生态减少工作量。先解决Python环境。如果你还没装好建议用Python 3.8以上版本安装时勾选“Add Python to PATH”否则后面pip命令会找不到。装完后在终端执行python -m pip install numpy pandas scikit-learn这三个库分别负责数值计算、表格处理和模型评估。装完后可以用python -c import numpy; print(numpy.__version__)验证。环境出问题时八成是环境变量没配好或者终端里用的是旧版本Python可以先试试python -m pip而不是裸pip。3.2 准备数据集用自造的“今天出门吗”还是鸢尾花手写决策树不适合一上来就用MNIST这种大样本。我建议先用一个经典小数据集根据天气、温度、湿度、风速决定要不要去打球。这是一个14条样本的表格特征全是离散值非常适合手算验证。import pandas as pd data pd.DataFrame({ outlook: [sunny,sunny,overcast,rainy,rainy,rainy, overcast,sunny,sunny,rainy,sunny,overcast, overcast,rainy], temperature: [hot,hot,hot,mild,cool,cool,cool, mild,cool,mild,mild,mild,hot,mild], humidity: [high,high,high,high,normal,normal,normal, high,normal,normal,normal,high,normal,high], windy: [false,true,false,false,false,true,true, false,false,false,true,true,false,true], play: [no,no,yes,yes,yes,no,yes,no,yes, yes,yes,yes,yes,no] }) print(data.groupby(play).size())这段代码直接构造DataFrame最后一列play是标签。输出显示yes有9个、no有5个说明类别不完全平衡但可接受。之所以用这种数据是因为每个特征只有两三个取值手工算一遍信息增益就能确认代码对不对。等你把三种算法跑通再换鸢尾花数据集也不迟。3.3 离散特征编码与连续值处理决策树和逻辑回归不一样不需要做OneHot编码。甚至做了OneHot反而坏事比如outlook原本三个取值OneHot会拆成三列每列只有0和1分裂时会变得非常碎。我们要做的只是区分“类别型特征”和“数值型特征”。numeric_cols data.select_dtypes(include[int64, float64]).columns.tolist() categorical_cols data.select_dtypes(include[object]).columns.tolist() print(numeric_cols, categorical_cols)这条代码用Pandas的dtype自动判断。打球数据集全是object所以numeric_cols为空。如果换到鸢尾花sepal length这种列就会进numeric_cols。后面实现C4.5和CART时对这两类特征要采用不同的分裂策略。3.4 定义树节点与递归出口建树前先定义一个TreeNode类后面三种算法都用它存树结构。class TreeNode: def __init__(self, featureNone, thresholdNone, labelNone): self.feature feature # 当前节点用于分裂的特征名 self.threshold threshold # 连续特征的分裂阈值离散特征为None self.subtree {} # 子节点字典CART里固定有left和right self.label label # 叶子节点的类别 self.fallback_label None # 预测时遇到未知取值时的默认类别递归出口有三个一是当前节点的样本全部属于同一类别直接生成叶子二是特征已经用完返回出现次数最多的类别三是分裂准则算出的增益或增益率、基尼指数增益没有改善同样返回多数类。这三个出口在三种算法里通用只是判断条件里那个“增益”替换成对应准则。4. 从零实现三种算法核心递归代码与参数全解析4.1 ID3信息增益计算与最佳特征选择ID3的框架是所有实现里最干净的。先写一个工具函数按特征值切分数据def split_by_categorical(data, feature, value): # 取feature列等于value的子集并删掉这一列 return data[data[feature] value].drop(columnsfeature)删掉这一列是因为决策树每个离散特征在路径上只用一次。接下来是选特征def choose_best_feature_id3(data, targetplay): base_entropy entropy(data[target]) best_gain -np.inf best_feature None for feature in data.columns.drop(target): values data[feature].unique() new_entropy 0.0 for v in values: subset data[data[feature] v] prob len(subset) / len(data) new_entropy prob * entropy(subset[target]) gain base_entropy - new_entropy if gain best_gain: best_gain gain best_feature feature return best_feature, best_gain函数的逻辑是按照每个离散特征的所有取值切分子集计算加权条件熵。参数target是标签列名默认play。注意这里用了data.columns.drop(target)确保不把标签列当作候选特征。返回的best_gain用于递归出口判断。建树主体def build_tree_id3(data, targetplay, max_depthNone, depth0): labels data[target] if len(np.unique(labels)) 1: node TreeNode(labellabels.iloc[0]) node.fallback_label labels.iloc[0] return node if len(data.columns.drop(target)) 0 or (max_depth and depth max_depth): majority labels.mode()[0] node TreeNode(labelmajority) node.fallback_label majority return node feature, gain choose_best_feature_id3(data, target) if gain 0: majority labels.mode()[0] node TreeNode(labelmajority) node.fallback_label majority return node node TreeNode(featurefeature) node.fallback_label labels.mode()[0] for value in data[feature].unique(): subset data[data[feature] value] if len(subset) 0: node.subtree[value] TreeNode(labellabels.mode()[0]) else: node.subtree[value] build_tree_id3(subset, target, max_depth, depth1) return node这里的max_depth是预剪枝参数depth是当前深度。很多人手写决策树时容易忘了限制深度导致树越建越深、过拟合严重。每个TreeNode都保存fallback_label这为后面处理预测时没见过的新取值留了后路。ID3整体就是这样递归、分叉、遇到纯节点就停。4.2 C4.5增益率、连续值阈值与缺失值处理C4.5比ID3麻烦在连续特征上。建树时不能删掉连续特征因为同一特征在不同深度用不同阈值可以再切。所以我们要把特征分成两类处理。先写连续特征的候选阈值搜索def choose_best_feature_c45(data, targetplay, numeric_colsNone): if numeric_cols is None: numeric_cols [] base_entropy entropy(data[target]) best_gain_ratio -np.inf best_feature None best_threshold None for feature in data.columns.drop(target): if feature in numeric_cols: values sorted(data[feature].unique()) for i in range(len(values) - 1): threshold (values[i] values[i 1]) / 2 left data[data[feature] threshold] right data[data[feature] threshold] if len(left) 0 or len(right) 0: continue new_entropy (len(left) / len(data) * entropy(left[target]) len(right) / len(data) * entropy(right[target])) gain base_entropy - new_entropy split_info - (len(left) / len(data) * np.log2(len(left) / len(data)) len(right) / len(data) * np.log2(len(right) / len(data))) if split_info 0: continue ratio gain / split_info if ratio best_gain_ratio: best_gain_ratio ratio best_feature feature best_threshold threshold else: values data[feature].unique() new_entropy 0.0 split_info 0.0 for v in values: subset data[data[feature] v] prob len(subset) / len(data) new_entropy prob * entropy(subset[target]) if prob 0: split_info - prob * np.log2(prob) if split_info 0: continue gain base_entropy - new_entropy ratio gain / split_info if ratio best_gain_ratio: best_gain_ratio ratio best_feature feature best_threshold None return best_feature, best_threshold这段代码把连续特征和离散特征分到同一个循环里返回值统一为(feature, threshold)其中离散特征的threshold为None。数值特征的核心是“排序后取相邻中点”注意values必须排序否则阈值算出来没有意义。候选阈值很多时计算会很慢可以只取分位数但在小数据集上直接全遍历即可。C4.5建树时不能像ID3那样丢掉特征列否则连续特征无法在下一层复用。所以递归函数里要传当前候选特征列表而不是简单dropdef build_tree_c45(data, targetplay, featuresNone, numeric_colsNone, max_depthNone, depth0): labels data[target] if len(np.unique(labels)) 1: node TreeNode(labellabels.iloc[0]) node.fallback_label labels.iloc[0] return node if features is None or len(features) 0 or (max_depth and depth max_depth): majority labels.mode()[0] node TreeNode(labelmajority) node.fallback_label majority return node # 这里省略了根据best_threshold切分的细节离散特征仍按值分连续特征按阈值分连续特征在递归时需要把已经用过的阈值移除否则同一个阈值可能被反复选中。比较省事的做法是每次分裂后把该特征从numeric_cols的候选里删掉。严格来说C4.5允许同一个特征用不同阈值多次分裂但为了控制过拟合我一般选择每个连续特征最多用一次。4.3 CART基尼指数与二叉树生成CART的实现要点是“二分一切”。离散特征也不再多叉而是把“等于某个值”分到左子树“不等于该值”分到右子树。连续特征和C4.5一样按阈值分。def choose_best_split_cart(data, targetplay, numeric_colsNone): if numeric_cols is None: numeric_cols [] best_gini np.inf best_feature None best_partition None for feature in data.columns.drop(target): if feature in numeric_cols: values sorted(data[feature].unique()) for i in range(len(values) - 1): threshold (values[i] values[i 1]) / 2 left data[data[feature] threshold] right data[data[feature] threshold] gini_idx (len(left) / len(data) * gini(left[target]) len(right) / len(data) * gini(right[target])) if gini_idx best_gini: best_gini gini_idx best_feature feature best_partition threshold else: for value in data[feature].unique(): left data[data[feature] value] right data[data[feature] ! value] gini_idx (len(left) / len(data) * gini(left[target]) len(right) / len(data) * gini(right[target])) # 离散特征二分的partition标记为取值本身 if gini_idx best_gini: best_gini gini_idx best_feature feature best_partition value return best_feature, best_partitionCART对于离散特征左子树是feature value右子树是feature ! value所以建树时左右两边用的切分条件不同。TreeNode的subtree字典里我统一用left和right两个键。这样predict_one时只需要判断样本特征和当前节点的threshold值。CART建树停止条件除了纯节点外还要加一个限制如果切分后的任一子集样本数小于min_samples_leaf就不切。这是最常用的预剪枝手段。def build_tree_cart(data, targetplay, numeric_colsNone, max_depthNone, min_samples_leaf1, depth0): labels data[target] if len(np.unique(labels)) 1: node TreeNode(labellabels.iloc[0]) node.fallback_label labels.iloc[0] return node if len(data) 2 * min_samples_leaf or (max_depth and depth max_depth): majority labels.mode()[0] node TreeNode(labelmajority) node.fallback_label majority return node feature, partition choose_best_split_cart(data, target, numeric_cols) if feature is None: majority labels.mode()[0] node TreeNode(labelmajority) node.fallback_label majority return node node TreeNode(featurefeature, thresholdpartition) node.fallback_label labels.mode()[0] if partition in numeric_cols: # 实际应判断feature的数据类型这里示意 left data[data[feature] partition] right data[data[feature] partition] else: left data[data[feature] partition] right data[data[feature] ! partition] node.subtree[left] build_tree_cart(left, target, numeric_cols, max_depth, min_samples_leaf, depth1) node.subtree[right] build_tree_cart(right, target, numeric_cols, max_depth, min_samples_leaf, depth1) return node注意上面if partition in numeric_cols只是示意实际要判断feature是否数值型。CART的min_samples_leaf是防止树末尾出现太小的子集这个方法在建树时非常好用。4.4 共用部分预测与统一封装三种算法建出的树结构不完全一致但预测逻辑可以统一。关键是处理CART二叉和ID3多叉的差异。def predict_one(node, sample): # sample: pandas Series一行样本 if node.label is not None: return node.label if node.threshold is not None: # CART节点threshold可能是数值阈值也可能是离散特征取值 if sample[node.feature] node.threshold: return predict_one(node.subtree[left], sample) else: return predict_one(node.subtree[right], sample) # ID3/C4.5多叉节点按特征值找子树 value sample[node.feature] if value in node.subtree: return predict_one(node.subtree[value], sample) else: # 新取值退回该节点多数类 return node.fallback_label这段代码最核心的地方是fallback_label的使用。很多手写决策树翻车都翻在这里训练集没有覆盖测试集的某个取值。设置fallback_label后至少不会报KeyError模型还能给出一个合理猜测。最后封装成一个统一类方便对比三种算法class DecisionTreeManual: def __init__(self, algorithmid3, max_depthNone, min_samples_leaf1): self.algorithm algorithm.lower() self.max_depth max_depth self.min_samples_leaf min_samples_leaf self.tree_ None def fit(self, X, y): data X.copy() data[__target__] y numeric_cols data.select_dtypes(include[int64, float64]).columns.tolist() numeric_cols.remove(__target__) if self.algorithm id3: self.tree_ build_tree_id3(data, target__target__, max_depthself.max_depth) elif self.algorithm c45: self.tree_ build_tree_c45(data, target__target__, numeric_colsnumeric_cols, max_depthself.max_depth) elif self.algorithm cart: self.tree_ build_tree_cart(data, target__target__, numeric_colsnumeric_cols, max_depthself.max_depth, min_samples_leafself.min_samples_leaf) else: raise ValueError(algorithm must be id3, c45 or cart) def predict(self, X): return X.apply(lambda row: predict_one(self.tree_, row), axis1)构造函数里的algorithm参数决定了走哪条建树路径。fit方法里把标签列暂时并进DataFrame这样三个建树函数可以共用相同的内部逻辑。调用时只需要model DecisionTreeManual(algorithmcart); model.fit(X_train, y_train); pred model.predict(X_test)和sklearn的接口长得一样。5. 手写决策树的五个坑从NaN到特征值缺失的排错记录5.1 现象信息增益计算遇到NaN程序直接中断原因出在entropy函数里。某个子集只有一个样本或全部类别一致时概率为0的类别算log2(0)会产生负无穷加总后变成NaN。继续递归整个节点全废。解决entropy函数里加一条probs probs[probs 0]只保留非零概率def entropy(labels): classes, counts np.unique(labels, return_countsTrue) probs counts / len(labels) probs probs[probs 0] # 去掉0概率 return -np.sum(probs * np.log2(probs))另外在切分子集时如果某个特征取值对应的子集为空直接给那个分支挂一个多数类叶子而不是继续递归。这个处理在build_tree里我已经写进去了。5.2 现象C4.5连续特征“上瘾”树越切越细训练集完美但测试集翻车连续特征候选阈值很多算法总能找到一个阈值把当前数据完美切开导致树深度失控。这本质上是过拟合。解决必须加min_samples_leaf。每次分裂前检查左右子集的样本数如果任何一侧小于min_samples_leaf就放弃这次分裂。我在CART的build_tree里已经加了len(data) 2 * min_samples_leaf的判断C4.5也要照做。实际使用中min_samples_leaf设置在2到5之间比较稳。5.3 现象CART输出树不是二叉树和预期不符有人把CART实现成了多叉离散特征有多少取值就分多少支。这样虽然也能建树但已经不是CART了后剪枝时会非常难处理sklearn的树也不长这样。解决强制二分。离散特征分裂成feature value和feature ! value两支而不是枚举所有取值。代码里choose_best_split_cart的离散特征分支就是这么做的。验证方法很简单打印递归深度每个内部节点的subtree字典应该只有left和right两个键。5.4 现象递归深度报错RecursionError: maximum recursion depth exceededPython默认递归上限是1000。如果数据噪声大、没有设置max_depth决策树会一直分裂直到每个叶子一个样本很快超过限制。解决在build_tree里增加depth参数达到max_depth就返回多数类叶子。如果实在需要很深的树可以在文件开头加import sys; sys.setrecursionlimit(10000)。但我的习惯是先调max_depth别动递归上限否则内存可能先爆。5.5 现象预测时KeyError测试样本的特征取值训练集没出现过这是离散特征最常见的坑。训练集里outlook只有sunny/overcast/rainy测试集冒出一个snowypredict_one里查subtree字典直接KeyError。解决TreeNode里保存fallback_label建树时把当前节点的多数类存进去。predict_one查不到子节点时返回fallback_label。这个设计在4.4节代码里已经体现了。另外在建树时给每个节点都手动赋值fallback_label别等预测时才发现属性不存在。6. 把三种手写树调成可用模型剪枝、参数与对比验证6.1 预剪枝参数怎么加手写树最容易被吐槽的就是只会无脑分裂。我一般保留下面这几个参数参数作用建议初始值max_depth限制树的深度5~8min_samples_split节点样本数小于此值不再分裂2~10min_samples_leaf叶子至少要有的样本数1~5min_impurity_decrease分裂前后不纯度下降的最小值0~0.01这些参数听起来和sklearn很像但手写实现里它们都要自己放进递归条件里。比如min_samples_split要在choose_best之前判断len(data) min_samples_split就停。参数不是越大越好可以用验证集过一遍选准确率最高的组合。6.2 后剪枝一种简单的代价复杂度剪枝思路预剪枝容易欠拟合后剪枝更稳。代价复杂度剪枝CCP的思路是对每个内部节点计算“剪掉子树替换成叶子后的误差增加量 α * 叶子数减少量”选α从小到大的顺序剪。手写时不需要完整复现sklearn的ccp_alpha搜索可以做一个简化版用训练集建出完整树自底向上遍历每个内部节点临时把它的子树替换成多数类叶子在验证集上计算准确率如果替换后准确率不下降就保留替换否则恢复原样。这个办法很直观代码量也不大。缺点是要在树结构上做“临时替换”再“还原”需要深拷贝或者额外标记实现时记得别改坏原始树。6.3 用验证集评估你的手写树准确率、混淆矩阵与sklearn对比写完三种算法我习惯直接用同一个数据集和sklearn对比这样能快速验出实现里有没有隐藏bug。from sklearn.tree import DecisionTreeClassifier from sklearn.metrics import accuracy_score, confusion_matrix # 假设已经把数据切好训练集、测试集 # X_train, X_test, y_train, y_test 已就绪 my_model DecisionTreeManual(algorithmcart, max_depth3) my_model.fit(X_train, y_train) my_pred my_model.predict(X_test) sk_model DecisionTreeClassifier(criteriongini, max_depth3) sk_model.fit(X_train, y_train) sk_pred sk_model.predict(X_test) print(手写CART准确率:, accuracy_score(y_test, my_pred)) print(sklearn准确率:, accuracy_score(y_test, sk_pred)) print(confusion_matrix(y_test, my_pred))如果准确率差距在0.05以内说明实现基本没问题。如果差距很大优先检查连续特征阈值计算和剪枝条件。sklearn的criterion参数entropy对应ID3/C4.5的思路gini对应CART。6.4 一个值得保留的习惯把树打印出来看我调试手写树时最受益的一个习惯是写一个print_tree函数把每个节点的特征、阈值、子节点缩进打印出来。def print_tree(node, indent): if node.label is not None: print(f{indent}叶子: {node.label}) return print(f{indent}特征: {node.feature}, 阈值: {node.threshold}) for key, child in node.subtree.items(): print(f{indent} [{key}]) print_tree(child, indent )打印出来之后你能直观看到树是否太深、哪些特征被反复使用、叶子样本是否太少。树是一棵能看懂的模型这也是决策树区别于神经网络的魅力。我自己最常犯的错是忘了处理空子集一debug就是半小时后来养成了每次写完分裂函数先用那个14条数据的打球数据集手工验算根节点选择再跑递归的习惯。这个习惯帮我省了大量调树时间。希望帮到你。本文还有配套的精品资源点击获取
返回列表