Python-sklearn-决策树 Sklearn 决策树sklearn.tree提供决策树分类和回归模型。 分类树DecisionTreeClassifier⭐fromsklearn.treeimportDecisionTreeClassifier,plot_tree modelDecisionTreeClassifier(criteriongini,# 分裂标准# gini / entropy / log_losssplitterbest,# best 或 randommax_depthNone,# 树最大深度None不限min_samples_split2,# 内部节点最小样本数min_samples_leaf1,# 叶节点最小样本数min_weight_fraction_leaf0.0,max_featuresNone,# 每次分裂考虑的特征数# None / int / float / sqrt / log2 / autorandom_state42,max_leaf_nodesNone,# 最大叶节点数min_impurity_decrease0.0,# 最小不纯度降低class_weightNone,# balanced / dict / Noneccp_alpha0.0# 最小代价复杂度剪枝参数)model.fit(X,y)# 关键属性print(model.feature_importances_)# 特征重要性print(model.classes_)# 类别数组print(model.n_classes_)# 类别数print(model.n_features_in_)# 特征数print(model.n_outputs_)# 输出数print(model.tree_)# 底层 Tree 对象# 树结构详细属性treemodel.tree_print(tree.node_count)# 节点总数print(tree.max_depth)# 树实际深度print(tree.n_leaves)# 叶节点数print(tree.children_left)# 左子节点索引数组print(tree.children_right)# 右子节点索引数组print(tree.feature)# 每个节点分裂的特征索引print(tree.threshold)# 每个节点分裂的阈值print(tree.value)# 每个节点的类别分布print(tree.impurity)# 每个节点的不纯度print(tree.n_node_samples)# 每个节点的样本数# 预测方法y_predmodel.predict(X)y_probmodel.predict_proba(X)# 各类别概率y_log_probmodel.predict_log_proba(X)# apply: 返回每个样本的叶节点索引leaf_indicesmodel.apply(X)# decision_path: 返回决策路径稀疏矩阵pathmodel.decision_path(X) 回归树DecisionTreeRegressor⭐fromsklearn.treeimportDecisionTreeRegressor modelDecisionTreeRegressor(criterionsquared_error,# 分裂标准# squared_error(MSE)# friedman_mse(含Friedman调整的MSE)# absolute_error(MAE)# poisson(泊松偏差)splitterbest,max_depthNone,min_samples_split2,min_samples_leaf1,min_weight_fraction_leaf0.0,max_featuresNone,random_state42,max_leaf_nodesNone,min_impurity_decrease0.0,ccp_alpha0.0)model.fit(X,y)y_predmodel.predict(X)leaf_indicesmodel.apply(X) 可视化决策树1.plot_tree()— Matplotlib 可视化 ⭐fromsklearn.treeimportplot_treeimportmatplotlib.pyplotasplt plt.figure(figsize(20,10))plot_tree(model,filledTrue,# 填充颜色反映类别分布roundedTrue,# 圆角节点fontsize10,feature_namesfeature_names,class_namesclass_names,proportionFalse,# True 显示比例而非绝对数impurityTrue,# 显示不纯度labelroot,# all,root,noneprecision3# 数值精度)plt.show()2.export_text()— 文本导出fromsklearn.treeimportexport_text textexport_text(model,feature_namesfeature_names,max_depth3,spacing3,decimals2,show_weightsFalse)print(text)输出示例:|--- feature_2 2.45 | |--- class: setosa |--- feature_2 2.45 | |--- feature_3 1.75 | | |--- class: versicolor ...3.export_graphviz()— Graphviz 导出fromsklearn.treeimportexport_graphvizimportgraphviz dot_dataexport_graphviz(model,out_fileNone,feature_namesfeature_names,class_namesclass_names,filledTrue,roundedTrue,special_charactersTrue)graphgraphviz.Source(dot_data)graph.render(decision_tree,formatpng)✂️ 剪枝决策树容易过拟合通过以下参数控制预剪枝Pre-pruning# 限制树的生长modelDecisionTreeClassifier(max_depth5,# 限制深度min_samples_split20,# 分裂所需最少样本min_samples_leaf10,# 叶节点最少样本max_leaf_nodes50,# 限制叶节点数量min_impurity_decrease0.01,# 不纯度降低阈值)后剪枝Post-pruning / CCPfromsklearn.treeimportDecisionTreeClassifier# 1. 先完整训练获取剪枝路径modelDecisionTreeClassifier(random_state42)pathmodel.cost_complexity_pruning_path(X_train,y_train)# 2. 查看不同 alpha 的影响alphaspath.ccp_alphas impuritiespath.impurities# 3. 用不同 alpha 训练并选择最佳models[]foralphainalphas:dtDecisionTreeClassifier(random_state42,ccp_alphaalpha)dt.fit(X_train,y_train)models.append(dt)# 4. 比较train_scores[m.score(X_train,y_train)forminmodels]test_scores[m.score(X_test,y_test)forminmodels] 特征重要性importnumpyasnpimportmatplotlib.pyplotaspltdefplot_feature_importance(model,feature_namesNone,top_n10):绘制特征重要性importancesmodel.feature_importances_ indicesnp.argsort(importances)[::-1][:top_n]iffeature_namesisNone:feature_names[fFeature{i}foriinrange(len(importances))]plt.figure(figsize(10,6))plt.barh(range(top_n),importances[indices],aligncenter)plt.yticks(range(top_n),[feature_names[i]foriinindices])plt.xlabel(Feature Importance)plt.gca().invert_yaxis()plt.title(Top Feature Importances)plt.tight_layout()plt.show() 调参指南防止过拟合的关键参数按优先级# 1. max_depth — 首先限制3~15 通常较好# 2. min_samples_split — 再限制分裂10~100# 3. min_samples_leaf — 限制叶节点5~50# 4. max_leaf_nodes — 直接限制复杂度# 5. ccp_alpha — 后剪枝modelDecisionTreeClassifier(max_depth8,min_samples_split20,min_samples_leaf10,max_leaf_nodes100,random_state42)常见问题问题原因解决过拟合树太深加大min_samples_split/min_samples_leaf减小max_depth欠拟合树太浅增加max_depth减小min_samples_split样本不均衡类别分布偏差设置class_weightbalanced特征过多噪音特征影响设置max_featuressqrtExtraTreeClassifier/ExtraTreeRegressor— 极端随机树与普通决策树不同分裂阈值完全随机。fromsklearn.treeimportExtraTreeClassifier,ExtraTreeRegressor modelExtraTreeClassifier(criteriongini,splitterrandom,# 必须为 randommax_depthNone,min_samples_split2,random_state42)model.fit(X,y)[[sklearn-总览|← 返回总览]] | [[sklearn-集成学习|集成学习 →]]

相关新闻

最新新闻

时序数据库TDengine在工业物联网中的核心优势与应用

时序数据库TDengine在工业物联网中的核心优势与应用

1. 为什么工业数据需要专门的时序数据库? 在工业物联网(IIoT)和智能制造领域,数据采集的典型特征与传统互联网应用有着本质区别。以某汽车工厂的传感器网络为例,2000多个测点以每秒10次的频率采集温度、振动、电流等数…

2026/8/11 11:55:19
问卷星在线考试vs垂直考试系统,2026中小企业培训选哪个?

问卷星在线考试vs垂直考试系统,2026中小企业培训选哪个?

1.导语对于中小企业的HR与培训岗从业者来说,选到适配的在线考试系统,能直接降低培训考核的时间与人力成本。本文结合问卷星、雨课堂等多款不同类型考试工具的实际使用体验,从功能、成本、适配性等维度展开对比,帮大家快速匹配适合…

2026/8/11 11:55:19
企业做专利检索,数据库到底该怎么选?

企业做专利检索,数据库到底该怎么选?

在很多人的印象里,专利检索要么是法务的事,要么是少数专利代理人的活儿。但真正带过研发团队、做过项目申报、跑过产品上市、参与过投融资尽调的人,感受不太一样——专利数据库早就不是小圈子的工具,它越来越像企业创新活动里的基…

2026/8/11 11:55:19
2025 黑马程序员 AI运维云计算 AI全程赋能,2025

2025 黑马程序员 AI运维云计算 AI全程赋能,2025

拥抱AI时代运维:黑马程序员AI运维云计算,AI全程赋能岗位实战 在云计算、微服务与分布式架构全面普及的今天,企业IT系统的复杂度呈指数级增长。一个线上故障可能涉及数十个服务、上百个节点、千万级调用链路,传统的“看监控、查日志…

2026/8/11 11:55:19
能源计量结算平台技术解析与行业实践

能源计量结算平台技术解析与行业实践

1. 项目背景与行业意义 宁夏宝丰集团作为西北地区重要的能源化工企业,其水电表计量结算管理平台的建设具有典型的行业示范价值。这类项目通常涉及能源计量、数据采集、费用结算等核心业务环节,是工业企业实现精细化管理和数字化转型的基础设施。 在传统…

2026/8/11 11:55:19
eShop本地开发环境配置与微服务部署指南

eShop本地开发环境配置与微服务部署指南

1. eShop本地运行环境准备作为电商系统开发者的标配工具,eShop的本地化部署能极大提升开发调试效率。不同于直接操作线上环境,本地运行可以自由测试支付回调、订单状态修改等敏感操作,还能避免多人协作时的环境冲突问题。我经手过三个不同技术…

2026/8/11 11:50:19