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-集成学习|集成学习 →]]