ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

回归树原理详解:从CART分裂逻辑到剪枝与实战

回归树原理详解:从CART分裂逻辑到剪枝与实战 回归树这玩意儿我在实际项目中用了不少次很多朋友一上来就看复杂公式结果绕晕了。其实它背后的逻辑特别朴素把数据切成几块每块里用最简单的平均值做预测。今天我不堆公式就拿它最直观的思路和实际代码/计算过程来拆一遍顺便把CART回归树的两个关键点——怎么切、怎么防过拟合——讲透。适合刚入门机器学习、或者已经用过线性回归想换个非线性方案的朋友参考。1. 回归树到底在做什么1.1 抛开公式先用生活场景理解假设你在一家奶茶店当店长想根据“当天最高温度”预测“奶茶销量”。你做了一张表记录了一整个月的数据32度的日子卖出180杯28度卖出160杯20度卖出130杯15度卖出90杯8度卖出60杯。这时候如果非要用一条直线去拟合“温度—销量”你会发现关系不是直线天热的时候销量增速放缓天冷的时候销量下降得也快。线性回归一条直线掰不过来。回归树的做法完全不同它会自动找到几个“分界温度”比如“22度”和“30度”把温度区间切成三段。然后每段里直接取一个平均值作为预测值。比如温度在8到22度之间平均销量按样本均值估计是75杯22到30度之间均值估计是145杯30度以上均值估计是178杯。树的结构就像在问“今天温度超过22度了吗”“超过了那超过30度了吗”每回答一个Yes/No就往下一层走最后落在叶子节点上叶子节点里存的就是这个区间的平均销量。这就是回归树最核心的思想特征空间划分成若干矩形区域每个区域对应一个叶子节点叶子节点的预测值就是该区域内训练样本目标变量的均值。你说它简单确实简单但它能拟合非常复杂的非线性关系这是线性模型做不到的。1.2 回归树和分类树的区别很多初学者把回归树和分类树混在一起其实核心区别就一个叶子节点输出什么。分类树的叶子节点输出一个类别标签比如“会买/不会买”或者输出一个概率分布回归树的叶子节点输出一个实数比如销量、房价、温度。这个差异直接决定了分裂时的评判标准也不一样。分类树爱用Gini系数或信息增益回归树用的是均方误差MSE或均方误差减少量。后面我会重点解释MSE是怎么驱动回归树生长的。另一个常见的误解是回归树既然叫“树”它是不是只能处理表格数据实际应用中处理表格型、结构化数据树的优势非常大。图像、文本、语音这类非结构化数据树模型不擅长那是深度学习的地盘。所以拿到一个项目先看数据类型结构化表格数据特征和目标变量关系复杂非线性有足够的样本量这类场景我非常推荐优先试回归树或基于树的集成模型。2. CART回归树的分裂逻辑拆解2.1 分裂标准为什么是MSE减少量CARTClassification And Regression Tree是Breiman等人在1984年提出的经典算法回归树部分的标准做法是每次分裂时选择一个特征和该特征的一个取值作为分裂点把当前节点的样本分成左右两份目标是让分裂后的两份数据各自内部的“纯度”更高。这里的纯度用均方误差衡量。假设当前节点有N个样本输出值y的均值是ŷ那么节点内的总平方误差是SSE Σ(yi - ŷ)²这个值反映了样本围绕均值波动的程度。如果均值差很不稳定SSE就大。我们希望分裂后左右两个子节点的SSE加起来尽量小。选分裂点时遍历所有特征、所有可能的分裂值计算分裂后的总SSE SSE_left SSE_right取总SSE最小的那个特征和值。有经验的朋友可能会问为什么不直接用每个子节点的MSE而用SSE因为SSE是带样本量权重的如果子节点样本太少MSE可能意外地小但SSE能规避这种虚假的“纯度”。你在用sklearn的DecisionTreeRegressor时参数criterion默认就是squared_error本质上就是最小化SSE。2.2 回归树的生长过程从根到叶一颗回归树从根节点开始生长每一步都在做同一件事枚举当前节点数据的所有特征假设特征已经预处理过无缺失值。对每个特征把它的取值排序枚举所有“相邻取值的中点”作为候选分裂点。对每个候选分裂点把样本分到左、右两个子节点计算分裂后的总SSE。选出总SSE最小的特征分裂点组合执行分裂。对子节点递归重复上述过程直到满足停止条件。这里的停止条件通常包括节点样本数小于min_samples_split、树的深度达到max_depth、或者分裂后总SSE减少量小于某个阈值。不用把停止条件设得太激进等后面讲剪枝时再细说实际上就算你不设任何限制树也能一直长到每个叶子只剩一个样本然后GPU上跑个十层八层就可能过拟合。所以停止条件本质就是在控制模型复杂度。2.3 数据如何分割和离散化处理CART默认只做二元分裂也就是说一个节点只能分出两个子节点不是多路分叉。这个设计有它的道理二元分裂和多路分裂相比不需要传递“分支个数”这一超参数训练时也不会有节点过度碎片化的问题。而且很多实际场景里特征本质上是可以反复参与分裂的同一个特征可以在不同层再次出现。比如刚才奶茶店的例子“温度”这一列可能在根节点用“22度”切了一次当温度小于22度时又用“15度”再切一次。这说明回归树能自动捕获特征在不同取值区间上不同的影响模式这是它表达非线性关系的重要途径。对于类别型特征CART的做法是把类别编码后当成数值处理但更稳妥的做法是先用序数编码或者one-hot编码。如果类别是有序的比如学历小学、初中、高中、大学可以直接用有序整数编码如果是无序类别比如城市、品牌one-hot更保险虽然会带来维度上升但树模型对稀疏特征的耐受度比较高。3. 一个完整的小样本手动推导3.1 数据准备和第一次分裂咱们手动算一个极小的例子彻底搞懂回归树每一步在干什么。假设我们有6个样本特征x和输出y如下样本xy11522733944755126614首先计算根节点的SSE。整体均值 (57971214)/6 9整体SSE (5-9)² (7-9)² (9-9)² (7-9)² (12-9)² (14-9)² 16 4 0 4 9 25 58。现在枚举x的所有分裂点。x排序后是1、2、3、4、5、6相邻中点为1.5、2.5、3.5、4.5、5.5。我们逐个算分裂点1.5左样本只有第1个y5均值5SSE左0右样本为第2到6个y7,9,7,12,14均值(7971214)/59.8SSE右(7-9.8)²(9-9.8)²(7-9.8)²(12-9.8)²(14-9.8)² 7.840.647.844.8417.64 38.8。总SSE38.8。分裂点2.5左为样本1和2均值(57)/26SSE左(5-6)²(7-6)²2右为样本3到6均值(971214)/410.5SSE右(9-10.5)²(7-10.5)²(12-10.5)²(14-10.5)² 2.2512.252.2512.2529。总SSE31。分裂点3.5左样本1、2、3均值7SSE左(5-7)²(7-7)²(9-7)²8右样本4、5、6均值11SSE右(7-11)²(12-11)²(14-11)²161926。总SSE34。分裂点4.5左样本1到4均值7SSE左8右样本5、6均值13SSE右112。总SSE10。分裂点5.5左样本1到5均值8SSE左(5-8)²(7-8)²(9-8)²(7-8)²(12-8)²91111628右样本6SSE右0。总SSE28。明显分裂点4.5的总SSE最小10所以第一步在x4.5处分裂。左分支包含样本1到4右分支包含样本5和6。3.2 继续生长直到停止左分支样本为(1,5)、(2,7)、(3,9)、(4,7)均值为7SSE8。对它继续枚举分裂点1.5、2.5、3.5。算一遍分裂点1.5SSE左0SSE右(7-7.67)²(9-7.67)²(7-7.67)²0.441.780.442.66总SSE2.66。分裂点2.5SSE左2SSE右(9-8)²(7-8)²2总SSE4。分裂点3.5SSE左(5-7)²(7-7)²(9-7)²8SSE右0总SSE8。所以左分支在x1.5处继续分裂。最终这棵小树长成x 1.5叶子预测值5。1.5 ≤ x 4.5叶子预测值7.67样本2、3、4的y均值。x ≥ 4.5叶子预测值13样本5、6的y均值。你看整棵树其实就是把x轴切成了三段每段给一个均值。这就是回归树最原始的样子。增加树的深度、增加特征数量本质上就是把这个“分段拟合”的过程变得更细、更复杂。3.3 为什么叶子用均值而不用更复杂的模型可能有人问既然每个区域里数据不一定是线性关系为什么不用一个线性模型来拟合每个叶子里的数据这个问题问到点子上了。把叶子节点里的模型换成线性回归就变成了“模型树”Model Tree。它的好处是每个区域内部利用线性关系做更精细的预测坏处是容易过拟合还要给每个叶子维护一套系数解释性变差。所以CART默认用均值简单、鲁棒、稳定。现实中即使是最经典的回归树实现也在这个细节上坚持用均值。只有在集成学习、或明确追求精度且样本量充足时我才会考虑在叶子内部再叠加线性模型。4. 剪枝防止回归树“背答案”4.1 什么是树的过拟合回归树天生容易过拟合。假设你不限制树的大小每个叶子节点分裂到只剩一个样本那训练集上的SSE会降到0看起来完美无误。但是遇到新数据这种“完美”就崩了。因为它已经把每个训练样本的具体数值背下来了而不是学到一个泛化模式。就像学生把练习册答案全部背下来考试一旦换题目就不会了。所以实际应用回归树时一定要控制复杂度。两个路径一是通过超参数限制树的生长预剪枝二是先让树长满再自底向上合并一些叶子后剪枝。sklearn里的DecisionTreeRegressor主要支持预剪枝R语言rpart包和CART原版算法则支持代价复杂度剪枝后剪枝。实际项目里我两种都会配合先用预剪枝控制一个合理的规模再用交叉验证选后剪枝的强度参数。4.2 代价复杂度剪枝CCP的逻辑后剪枝的理论基础是代价复杂度剪枝。定义树的代价复杂度为Cα(T) 总SSE α * 叶子节点数这里的α是一个非负参数。第一项衡量拟合误差第二项衡量模型复杂度叶子越多模型越复杂。α越大对叶子节点数量的惩罚越大。剪枝的过程就是把那些“增加复杂度但并没有显著降低SSE”的子树剪掉直到在当前α下代价复杂度最小。实际操作时我们会先让树完全生长然后自底向上计算每个内部节点如果被剪掉、变成叶子节点代价复杂度的变化量。变化最小的节点先被剪掉。这样可以得到一系列不同规模的子树再用交叉验证从中选出泛化误差最小的那棵。sklearn中对应的是ccp_alpha参数你可以在DecisionTreeRegressor里设置ccp_alpha来控制后剪枝强度然后用GridSearchCV搜索合适的值。我这里给个实战建议用交叉验证找ccp_alpha时先在一个较宽的范围内搜索比如0到0.05观察树的大小和验证集误差的变化曲线。刚开始加大α验证集误差会下降因为剪掉了噪声细节继续加大α验证集误差会重新上升因为树变得太简单、欠拟合。曲线最低点对应的α就是当前数据集比较理想的值。4.3 最小叶子节点数和深度限制除了剪枝两个常用的预剪枝超参数是min_samples_leaf和max_depth。min_samples_leaf的意思是一个叶子节点至少要有多少样本。把min_samples_leaf设为5到20可以避免树生成太多只覆盖一两个样本的叶子。max_depth限制树的层数3到7层大多数场景下都够用。深度过深不仅过拟合还会让树的解释性大幅下降——你画出一棵20层的树根本没法跟业务方讲清楚。我在实际建模时的一般流程是先不设剪枝条件把决策树画出来看看它长到什么程度会开始过拟合接着依次调max_depth、min_samples_leaf、min_samples_split最后如果还嫌过拟合就上ccp_alpha后剪枝。每一步都用交叉验证评估不要只看训练集误差。5. 回归树的优势和局限5.1 相比线性回归的独特价值拿回归树和线性回归对比不是要分个高下而是看清各自适用场景。线性回归假设目标变量和特征之间是线性关系或可以通过特征变换变成线性关系而且对特征之间多重共线性敏感对异常值敏感。回归树没有这些假设。它能自动关注特征的交互作用比如“年龄大于30且收入大于20万”这类条件组合线性回归需要你手动构造交叉特征树模型天然就能捕捉。反过来线性回归的优势是可解释性强回归系数就是每个特征对目标变量的边际影响这在银行风控和医疗等需要合规解释的领域是硬需求。回归树虽然也能通过特征重要性来评估贡献但它毕竟是一个分段函数的组合解释起来比一个线性公式要费力。业务上如果必须给出“每个变量的具体影响大小”线性回归或带L1惩罚的线性模型常常更合适。5.2 特征重要性的解读陷阱回归树可以提供特征重要性但很多新手容易掉进一个坑sklearn里回归树的feature_importances_是基于节点分裂时SSE减少总量的加权和。这意味着一个特征如果在树的上层被选作分裂点重要性天然偏高如果两个特征高度相关重要性评分可能被分散导致你低估某个关键因素。在实际项目中我从不只看树模型输出的特征重要性。我会交叉验证一下单独用某个特征训练一棵树看预测效果然后把这个特征去掉看验证集误差上升多少。如果误差上升明显说明它确实重要。这种“置换重要性”的方法虽然简单但比直接读feature_importances_更可靠。5.3 外推能力差的背后原因回归树的预测值是叶子节点内训练样本的均值所以它永远不“超出”训练数据的范围。线性回归可以外推即使没有高昂收入对应的样本只要斜率合理也能预测一个极高收入对应的房价。回归树做不到。你给它一个收入500万的样本如果训练集里最大收入是100万它最多只能落在“收入大于50万”那个叶子节点里预测值也就在那个区间均值的水平。这既是缺点也是优点。在金融风控这类场景里尽量不要让模型做超出训练分布的外推预测树的这种保守性反而能避免一些离谱的预测。但在销售预测、增长率预估这类需要外推的场景你要么改用线性模型要么做特征变换要么使用树模型时心里清楚它的预测上限受限于训练数据范围。6. 实操中的几个关键细节6.1 用sklearn快速构建一个回归树下面用一段代码演示最基础的回归树训练过程数据集用sklearn自带的加利福尼亚房价数据代码非常短。import numpy as np import matplotlib.pyplot as plt from sklearn.datasets import fetch_california_housing from sklearn.model_selection import train_test_split, cross_val_score, GridSearchCV from sklearn.tree import DecisionTreeRegressor, plot_tree data fetch_california_housing() X, y data.data, data.target X_train, X_test, y_train, y_test train_test_split( X, y, test_size0.2, random_state42 ) reg DecisionTreeRegressor(max_depth4, min_samples_leaf10, random_state42) reg.fit(X_train, y_train) print(Train R2:, reg.score(X_train, y_train)) print(Test R2:, reg.score(X_test, y_test))这段代码里max_depth4是为了控制树不过深min_samples_leaf10保证每个叶子的样本数不会太少。输出R2后你会发现训练集和测试集差距不会太大这就是预剪枝起的作用。如果你把max_depth去掉训练R2会飙升到接近1测试R2反而会下降这就是过拟合的直接证据。6.2 网格搜索选择剪枝参数这里提供一个网格搜索ccp_alpha的实用代码片段帮助你自己根据数据集挑合适的剪枝强度。reg DecisionTreeRegressor(random_state42) path reg.cost_complexity_pruning_path(X_train, y_train) ccp_alphas, impurities path.ccp_alphas, path.impurities # 去掉最大值因为alpha无穷大时树只剩根节点没意义 for alpha in ccp_alphas: reg DecisionTreeRegressor(random_state42, ccp_alphaalpha) scores cross_val_score(reg, X_train, y_train, cv5, scoringr2) print(falpha{alpha:.4f}, CV R2{scores.mean():.4f})cost_complexity_pruning_path会返回一组候选alpha值你遍历它们并做交叉验证选平均R2最高的alpha。这个流程比你手动猜max_depth要稳得多。不过要注意ccp_alphas的数量可能很大实际使用中可以先粗筛一下每隔几个取一个值减少计算量。6.3 可视化回归树给业务方看回归树的一个巨大优势是可视化。你可以用plot_tree把树画成流程图直接给业务方看“如果一个客户的收入小于5万且年龄小于30预测消费金额是500元。”这种可读性是随机森林、XGBoost给不了的。但默认画出来的树节点上显示的信息太多字体小、不好看建议加参数。plt.figure(figsize(20, 10)) plot_tree( reg, feature_namesdata.feature_names, filledTrue, roundedTrue, fontsize10, max_depth3 ) plt.show()max_depth参数在plot_tree里只影响显示不会改变模型本身。画图时限制显示深度是为了视觉上清爽。如果你想把树保存成图片给业务方用plt.savefig(tree.png, dpi300)输出高清图。7. 回归树的典型应用场景和扩展方向7.1 单棵树不够用时怎么办如果数据量不大、特征关系比较线性单棵回归树往往够用。但真实场景中单棵树的预测精度通常拼不过集成模型。随着数据量增加、特征之间交互更复杂单棵树要么欠拟合限制太严要么过拟合限制太松很难找到恰到好处的平衡点。这时候就该上随机森林、梯度提升树了。随机森林的原理是训练多棵回归树每棵树用不同的自助采样子集和随机特征子集最后把预测结果平均。它通过“集体智慧”大幅降低了单棵树的方差稳定性非常好。梯度提升树则不同每棵树在前一棵树的残差上拟合通过逐步减少残差来逼近目标对异常值比较敏感但精度上限高。现在工业界大规模使用的XGBoost、LightGBM、CatBoost都是梯度提升树的不同工程实现。7.2 在数据挖掘流程中回归树承担什么角色在完整的数据挖掘流程中回归树可以是最终模型也可以是一个探索工具。我最常用的方式之一是先用深度较小的回归树做特征筛选和变量关系探查。比如页面转化率的预测中我先用depth3的树去看哪些渠道、哪些时段对转化影响最大。树把样本划分之后每个叶子区间内目标变量的均值变化能直观告诉我哪个群体表现好、哪个群体表现差。这种“分段分析”在业务诊断中比一堆假设检验更实用。回归树做缺失值填补也值得一试。把缺失的目标变量当预测目标其他特征齐全的样本做训练集训练一棵回归树然后用它预测缺失值。虽然比不上专门的插补方法比如MICE但胜在简单、无需分布假设、对非线性关系友好。7.3 和线性模型打组合拳分段线性化前面说了回归树是分段常数函数外推能力差。但如果先用回归树把样本划分成不同区间再在每个区间内拟合一个线性回归模型就能改善外推且保留部分可解释性。这就是带格子的线性组合模型业界也做过类似方案。实操上很简单用回归树深度3-5得到每个训练样本所属的叶子节点把叶子节点编号当成分组变量然后在每个组内分别训练线性回归。预测新样本时先用树决定它属于哪个组再用该组的线性模型预测。这种方法在一些业务场景里效果很好既能捕捉数据整体的非线性结构又能保留梯度信息用于外推。缺点是流程变复杂了如果叶子数量太多每个组内样本量不足线性模型容易不稳定。折中方案是让叶子组数量控制在5到10个之间。8. 实操中踩过的坑和排查技巧8.1 数据中的异常值影响回归树对异常值比线性回归要稳健得多因为它用均值预测单个异常值只会影响一个叶子节点的均值。但这个稳健性是相对的。如果你的某个叶子节点里恰好只有3个样本其中一个样本是异常值那预测结果就可能被带偏。我建议数据预处理时依然要粗略检查异常值尤其是目标变量里的极端值。如果业务逻辑上这些极端值有意义可以先保留如果只是噪声在训练前去掉或做缩尾处理效果更好。8.2 特征量纲和归一化的必要性回归树分裂只看特征取值的大小顺序不关心单位所以标准化或归一化不影响树的训练结果。这跟线性回归、KNN、SVM很不一样。你在实际建模时不用先做标准化省一步是一步。不过如果你后续要把回归树的结果当成特征输入到其他模型比如逻辑回归那时候才需要对树的叶子做编码或者对特征做归一化。8.3 类别不平衡怎么办回归问题的“不平衡”概念不同于分类问题但如果目标变量分布严重偏斜比如90%的样本集中在很小的区间少数样本极大回归树会倾向于把预测值往中位数附近拉。这种情况下可以考虑预测目标做log变换或者用分位数回归的思路——把目标变量按分位数映射到更均匀的空间训练后再逆变换回来。我自己在房价预测中遇到过类似问题对目标变量取对数后树模型的验证集误差明显下降。8.4 对树模型预测结果做平滑回归树的预测是阶梯函数预测结果是一段一段的常数在特征空间上看起来不够平滑。如果你追求预测结果平滑一些可以用随机森林替代单棵树因为随机森林把多棵树的阶梯结果平均了平滑度好很多。或者对预测结果使用后处理平滑比如KNN平滑、核平滑不过这会增加线上预测的复杂度。业务上如果不是特别在意平滑性其实单棵树的阶梯预测完全能接受。8.5 常见问题速查表问题可能原因排查/解决办法训练R2很高测试R2很低过拟合调小max_depth调大min_samples_leaf尝试ccp_alpha训练R2和测试R2都低欠拟合或特征无效检查特征质量增加特征减少预剪枝强度特征重要性过于集中在少数特征特征相关性强或树深度太浅做特征相关性分析用置换重要性做交叉验证预测值看起来“没变化”树过小或叶子节点均值接近增大树容量检查目标变量分布外推预测异常保守树模型天然不能外推换线性模型或在叶子内部用线性拟合9. 最后分享一个小经验我自己用回归树最多的地方反而不是直接拿它当最终模型而是拿它做“快速摸底”。接到一个新数据集先随便跑一棵浅树看看主要特征的作用方向和交叉效果几分钟就能对数据有个直觉后面再用更复杂的模型去调精度。树的这种“快速启发式”价值很容易被低估尤其在时间紧、业务场景不明朗的时候画一棵树比跑十轮特征工程快多了。如果你手头正好有个回归问题特征不算太多、也有解释需求建议直接先训练一棵深度4-5的回归树画出来看看。观察它分成几个区间、每个区间的均值差异、哪些特征留到了最后——这一步做完你对数据的理解会上一个台阶。之后再上手随机森林或梯度提升方向感会明确很多。
返回列表