ARTICLE DETAIL

资讯详情

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

机器学习入门:随机森林(Random Forest)

机器学习入门:随机森林(Random Forest) 机器学习入门随机森林Random Forest前言本文是机器学习入门系列的第五站。在上一站中我们深入学习了决策树明白了它是一棵通过如果……那么……规则进行分类的白盒模型。但单棵决策树往往有一个致命的弱点——容易过拟合且对数据的扰动非常敏感。那么如果我们种下一片森林让多棵树一起投票表决呢这就是今天的主角随机森林Random Forest。它汇聚众树之所长大幅提升了模型的准确率和稳定性。目录一、从单棵树到一片森林二、随机森林的核心原理三、随机森林 vs 决策树四、随机森林的优缺点五、典型应用场景六、随机森林核心 API 速查七、实战案例八、总结一、从单棵树到一片森林1.1 为什么要集成想象你要诊断一种疑难杂症场景A只咨询一位医生可能因个人状态或经验局限而误判。场景B邀请100位专家医生独立会诊然后投票决定结果。显然多数投票的群体决策比单一专家更稳妥。这就是集成的核心思想。1.2 什么是集成学习集成学习Ensemble Learning的核心逻辑很简单三个臭皮匠顶个诸葛亮。它组合多个弱学习器来构建一个强学习器。单个弱学习器效果一般但组合后长短互补、优劣相抵最终实现整体远胜部分的效果。随机森林正是集成学习最具代表性的算法之一——它以决策树为弱学习器让每棵树各有所长最终通过投票汇聚众智。1.3 为什么要放弃单棵树单棵决策树有两个致命的痛点对数据极度敏感数据微小变化就可能导致树结构截然不同稳定性差。过拟合风险极高倾向于死记硬背训练样本泛化能力弱。随机森林正是为此而生——通过双重随机化抵消单棵树的偏误用群体决策取代个人判断。二、随机森林的核心原理随机森林Random Forest是由多棵决策树组成的分类或回归模型。其核心思想可以概括为两个随机2.1 随机性一样本随机Bagging 思想传统决策树训练时使用的是全部训练样本。而在随机森林中每棵树在训练时并不使用所有数据。做法从总训练集中采用**有放回采样Bootstrap Sampling**的方式随机抽取一定数量的样本构成该树的专属训练集。效果这就导致每棵树看到的训练数据都不尽相同。有些样本可能被一棵树多次抽中而有些样本可能从未被抽到。2.2 随机性二特征随机在传统决策树中选择分裂节点时会遍历所有特征去寻找最优切分点。而在随机森林中为了避免所有树都长得过于相似引入了第二重随机做法在每棵树进行节点分裂时算法不会考察全部特征而是从总特征中随机抽取部分特征通常取总特征数的平方根仅在这部分特征中寻找最优切分。效果这确保了森林中的每棵树都各有侧重、相互独立从而降低整体方差。2.3 最终决策当输入一个新的测试样本时森林中所有的决策树都会给出自己的预测结果分类任务采用多数投票法Majority Voting——哪个类别得票最多就作为最终分类结果。回归任务采用平均值法Averaging——将所有树的预测结果取平均作为最终输出。2.4 特征重要性计算方式随机森林通过记录每个特征在所有树中参与分裂时所带来的不纯度减少量Gini 不纯度减少或信息增益并将这些减少量按特征汇总、归一化最终得到每个特征的重要性得分。得分越高说明该特征对分类或回归的贡献越大。这使得随机森林自带特征选择能力在实际项目中非常实用。三、随机森林 vs 决策树对比维度单棵决策树CART随机森林Random Forest模型复杂度低单模型高多棵树的集成过拟合风险极高容易死记硬背显著降低双重随机性有效抑制方差训练速度快较慢需训练多棵树但可并行特征重要性可通过 Gini 或信息增益计算内置特征重要性评估非常实用数据敏感度对数据扰动非常敏感极度稳定抗噪声能力强是否需要标准化否否同样基于树的阈值分裂规则可解释性强白盒模型可直观展示决策路径弱数百棵树投票难以直观理解四、随机森林的优缺点4.1 优点抗过拟合能力突出双重随机性使得模型即使在特征多、样本相对较少的情况下也不易过拟合。高维数据处理能力强无需提前做特征选择它能自动评估特征的重要性。能处理缺失值内置了缺失值处理机制通过代理分裂等方式。易于并行化每棵树相互独立非常适合在多核 CPU 上并行训练效率较高相对其他集成算法而言。对异常值不敏感基于树的模型天然对异常值有较强的鲁棒性。4.2 缺点可解释性较差决策树是白盒模型但当数百棵树一起决策时人类难以直观理解其内部逻辑。它可被视为一个黑盒中的白盒。在某些高噪声数据集上可能过拟合当数据噪声较大时随机森林仍可能过度学习噪声。模型体积较大需保存全部树的结构内存占用和推理时间相对较高。对稀疏数据表现一般在极度稀疏的高维数据如文本分类上随机森林的表现通常不如线性模型或 SVM。五、典型应用场景金融风控信用评分判断用户是否具有违约风险准确率通常优于单一的逻辑回归和决策树。特征工程与特征筛选利用随机森林输出的特征重要性剔除无关冗余特征为其他模型做降维铺垫。遥感与生态学处理多波段遥感影像进行土地覆盖分类和植被类型识别。医疗诊断辅助基于多个临床指标对疾病风险进行综合预测。推荐系统作为排序模型或点击率预测的基模型之一。六、随机森林核心 API 速查6.1 导包方式分类任务与回归任务分别从sklearn.ensemble中导入对应的类分类RandomForestClassifier回归RandomForestRegressor其他常配套使用的模块包括数据集划分train_test_split、交叉验证cross_val_score以及各类评估指标分类报告、混淆矩阵、MSE、R² 等。6.2 核心参数详解随机森林的参数可分为三类树结构参数、随机性参数、工程优化参数。参数名类型默认值说明树结构参数n_estimatorsint100森林中树的数量越多模型越稳定但训练和推理越慢。通常 100~300 即可达到较好效果max_depthint / NoneNone每棵树的最大深度默认 None 表示不限制树会完全生长推荐手动设置如 10~30以防过拟合max_leaf_nodesint / NoneNone树的最大叶节点数限制叶节点数量可有效控制模型复杂度防止过拟合min_samples_splitint / float2内部节点再分裂所需的最小样本数值越大树越保守过拟合风险越低min_samples_leafint / float1叶节点所需的最小样本数值越大树越平滑max_featuresint / str / floatsqrt每次分裂时随机考虑的特征数分类推荐sqrt即总特征数的平方根回归推荐总特征数的 1/3随机性参数bootstrapboolTrue是否采用有放回抽样True 即为 Bagging 方式False 则每棵树使用全部样本易过拟合oob_scoreboolFalse是否使用袋外样本计算验证分数设为 True 可省去额外划分验证集random_stateint / NoneNone随机种子固定后结果可复现工程优化参数n_jobsintNone并行训练的 CPU 核心数设为-1表示使用全部核心verboseint0训练过程日志输出级别0 不输出1 输出进度条warm_startboolFalse是否复用前一次训练结果增量训练适合在逐步增加n_estimators时使用6.3 常用属性训练完成后可通过以下属性获取模型内部信息属性名说明feature_importances_特征重要性最常用返回长度为特征数的数组值越大表示该特征对预测越关键oob_score_袋外评分需事先设置oob_scoreTrue可作为模型泛化能力的参考指标estimators_森林中所有决策树对象的列表可单独查看每棵树的结构n_features_in_训练时使用的特征数量classes_分类任务中所有类别的标签分类模型专属n_classes_分类任务中的类别数量分类模型专属6.4 常用方法方法名适用任务说明fit(X, y)分类 / 回归训练模型一切开始的地方predict(X)分类 / 回归预测新样本的类别分类或数值回归predict_proba(X)分类专属预测新样本属于每个类别的概率输出如[[0.1, 0.9]]表示 10% 概率为类别 090% 为类别 1predict_log_proba(X)分类专属概率的对数形式数值更稳定适合概率连乘场景score(X, y)分类 / 回归返回评估分数分类返回准确率回归返回决定系数R²apply(X)分类 / 回归返回每个样本在每棵树中落入的叶节点索引可用于理解样本在森林中的路径七、实战案例垃圾邮件识别7.1 案例背景垃圾邮件识别是机器学习在文本处理领域的经典应用。本案例使用的 spambase 数据集共 4601 条样本57 个特征构建随机森林分类器实现邮件的自动分类。7.2 完整代码importmatplotlib.pyplotaspltimportpandasaspdimportnumpyasnpfromsklearn.model_selectionimporttrain_test_split,cross_val_scorefromsklearn.ensembleimportRandomForestClassifierfromsklearnimportmetricsdefcm_plot(yt,yp):fromsklearn.metricsimportconfusion_matriximportmatplotlib.pyplotasplt cmconfusion_matrix(yt,yp)plt.matshow(cm,cmapplt.cm.Reds)plt.colorbar()forx_inrange(len(cm)):fory_inrange(len(cm)):plt.annotate(cm[x_,y_],xy(y_,x_),vacenter,hacenter)plt.ylabel(True label)plt.xlabel(Predicted label)returnplt# 导入数据dataspd.read_csv(spambase.csv)Xdatas.iloc[:,:-1]ydatas.iloc[:,-1]X_train,X_test,y_train,y_testtrain_test_split(X,y,test_size0.2,random_state42)# 交叉验证选择最优深度depth_rangerange(1,30)best_score-1best_depthNoneprint( 交叉验证11折)fordepthindepth_range:rfRandomForestClassifier(n_estimators100,max_depthdepth,random_state42,n_jobs-2)scorescross_val_score(rf,X_train,y_train,cv11,scoringrecall)score_meannp.mean(scores)print(fdepth{depth:2d}, recall{score_mean:.4f})ifscore_meanbest_score:best_scorescore_mean best_depthdepthprint(f\n最优深度:{best_depth}, 最优召回率:{best_score:.4f})# 使用最优深度训练模型rfRandomForestClassifier(n_estimators100,max_depthbest_depth,random_state42,n_jobs-2)rf.fit(X_train,y_train)# 训练集评估train_predrf.predict(X_train)print(\n训练集分类报告)print(metrics.classification_report(y_train,train_pred))cm_plot(y_train,train_pred).show()# 测试集评估test_predrf.predict(X_test)print(\n测试集分类报告)print(metrics.classification_report(y_test,test_pred))cm_plot(y_test,test_pred).show()# 特征重要性分析importancesrf.feature_importances_ impd.DataFrame(importances,columns[importances])closdatas.columns clos_1clos.values clos_2clos_1.tolist()closclos_2[0:-1]im[clos]clos imim.sort_values(by[importances],ascendingFalse)[:10]indexrange(len(im))plt.yticks(index,im.clos)plt.barh(index,im[importances])plt.show()输出示例 交叉验证11折 depth 1, recall0.5680 depth 2, recall0.7498 depth 3, recall0.8027 depth 4, recall0.8330 ······ depth25, recall0.9225 depth26, recall0.9246 depth27, recall0.9211 depth28, recall0.9239 depth29, recall0.9183 最优深度: 23, 最优召回率: 0.9274 训练集分类报告 precision recall f1-score support 0 1.00 1.00 1.00 2258 1 1.00 0.99 1.00 1419 accuracy 1.00 3677 macro avg 1.00 1.00 1.00 3677 weighted avg 1.00 1.00 1.00 3677 测试集分类报告 precision recall f1-score support 0 0.94 0.98 0.96 527 1 0.97 0.92 0.94 393 accuracy 0.95 920 macro avg 0.95 0.95 0.95 920 weighted avg 0.95 0.95 0.95 9207.3 关键步骤说明步骤说明数据加载Spambase 数据集共 4601 条样本57 个特征最后一列为标签1 表示垃圾邮件0 表示正常邮件交叉验证选参遍历深度 1~29采用 11 折交叉验证以召回率recall为评估指标选出最优深度模型训练固定最优深度树数量设为 100其余参数保持默认模型评估分别输出训练集和测试集的分类报告并通过混淆矩阵可视化预测结果特征重要性提取 Top 10 重要特征并可视化帮助理解哪些邮件特征最能区分垃圾邮件八、总结核心知识点速查知识点关键概念集成学习组合多个弱学习器构建强学习器Bagging有放回抽样构建训练集特征随机分裂时仅考虑部分特征最终决策多数投票分类/ 平均回归特征重要性基于 Gini 不纯度减少量汇总核心 API 一览用途对应模块 / 方法分类模型sklearn.ensemble.RandomForestClassifier回归模型sklearn.ensemble.RandomForestRegressor训练fit(X, y)预测predict(X)概率预测分类predict_proba(X)特征重要性feature_importances_注意事项要点说明无需剪枝双重随机性替代了剪枝操作树的数量并非越多越好100~300 后收益递减单机瓶颈树过多时内存和推理时间上升随机种子固定random_state确保结果可复现特征数选择分类用sqrt回归用1/3总特征数系列直达上篇机器学习入门决策树Decision Tree本篇机器学习入门随机森林Random Forest本文下篇机器学习入门朴素贝叶斯Naive Bayes
返回列表