Python-sklearn-决策树
2026/9/15 5:22:31 网站建设 项目流程

Sklearn 决策树

sklearn.tree提供决策树分类和回归模型。


🌳 分类树

DecisionTreeClassifier

fromsklearn.treeimportDecisionTreeClassifier,plot_tree model=DecisionTreeClassifier(criterion='gini',# 分裂标准# 'gini' / 'entropy' / 'log_loss'splitter='best',# 'best' 或 'random'max_depth=None,# 树最大深度(None=不限)min_samples_split=2,# 内部节点最小样本数min_samples_leaf=1,# 叶节点最小样本数min_weight_fraction_leaf=0.0,max_features=None,# 每次分裂考虑的特征数# None / int / float / 'sqrt' / 'log2' / 'auto'random_state=42,max_leaf_nodes=None,# 最大叶节点数min_impurity_decrease=0.0,# 最小不纯度降低class_weight=None,# 'balanced' / dict / Noneccp_alpha=0.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 对象# 树结构详细属性tree=model.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_pred=model.predict(X)y_prob=model.predict_proba(X)# 各类别概率y_log_prob=model.predict_log_proba(X)# apply: 返回每个样本的叶节点索引leaf_indices=model.apply(X)# decision_path: 返回决策路径(稀疏矩阵)path=model.decision_path(X)

📈 回归树

DecisionTreeRegressor

fromsklearn.treeimportDecisionTreeRegressor model=DecisionTreeRegressor(criterion='squared_error',# 分裂标准# 'squared_error'(MSE)# 'friedman_mse'(含Friedman调整的MSE)# 'absolute_error'(MAE)# 'poisson'(泊松偏差)splitter='best',max_depth=None,min_samples_split=2,min_samples_leaf=1,min_weight_fraction_leaf=0.0,max_features=None,random_state=42,max_leaf_nodes=None,min_impurity_decrease=0.0,ccp_alpha=0.0)model.fit(X,y)y_pred=model.predict(X)leaf_indices=model.apply(X)

🎨 可视化决策树

1.plot_tree()— Matplotlib 可视化 ⭐

fromsklearn.treeimportplot_treeimportmatplotlib.pyplotasplt plt.figure(figsize=(20,10))plot_tree(model,filled=True,# 填充颜色(反映类别分布)rounded=True,# 圆角节点fontsize=10,feature_names=feature_names,class_names=class_names,proportion=False,# True 显示比例而非绝对数impurity=True,# 显示不纯度label='root',# 'all','root','none'precision=3# 数值精度)plt.show()

2.export_text()— 文本导出

fromsklearn.treeimportexport_text text=export_text(model,feature_names=feature_names,max_depth=3,spacing=3,decimals=2,show_weights=False)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_data=export_graphviz(model,out_file=None,feature_names=feature_names,class_names=class_names,filled=True,rounded=True,special_characters=True)graph=graphviz.Source(dot_data)graph.render('decision_tree',format='png')

✂️ 剪枝

决策树容易过拟合,通过以下参数控制:

预剪枝(Pre-pruning)

# 限制树的生长model=DecisionTreeClassifier(max_depth=5,# 限制深度min_samples_split=20,# 分裂所需最少样本min_samples_leaf=10,# 叶节点最少样本max_leaf_nodes=50,# 限制叶节点数量min_impurity_decrease=0.01,# 不纯度降低阈值)

后剪枝(Post-pruning / CCP)

fromsklearn.treeimportDecisionTreeClassifier# 1. 先完整训练,获取剪枝路径model=DecisionTreeClassifier(random_state=42)path=model.cost_complexity_pruning_path(X_train,y_train)# 2. 查看不同 alpha 的影响alphas=path.ccp_alphas impurities=path.impurities# 3. 用不同 alpha 训练并选择最佳models=[]foralphainalphas:dt=DecisionTreeClassifier(random_state=42,ccp_alpha=alpha)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_names=None,top_n=10):"""绘制特征重要性"""importances=model.feature_importances_ indices=np.argsort(importances)[::-1][:top_n]iffeature_namesisNone:feature_names=[f'Feature{i}'foriinrange(len(importances))]plt.figure(figsize=(10,6))plt.barh(range(top_n),importances[indices],align='center')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 — 后剪枝model=DecisionTreeClassifier(max_depth=8,min_samples_split=20,min_samples_leaf=10,max_leaf_nodes=100,random_state=42)

常见问题

问题原因解决
过拟合树太深加大min_samples_split/min_samples_leaf,减小max_depth
欠拟合树太浅增加max_depth,减小min_samples_split
样本不均衡类别分布偏差设置class_weight='balanced'
特征过多噪音特征影响设置max_features='sqrt'

ExtraTreeClassifier/ExtraTreeRegressor— 极端随机树

与普通决策树不同,分裂阈值完全随机。

fromsklearn.treeimportExtraTreeClassifier,ExtraTreeRegressor model=ExtraTreeClassifier(criterion='gini',splitter='random',# 必须为 'random'max_depth=None,min_samples_split=2,random_state=42)model.fit(X,y)

[[sklearn-总览|← 返回总览]] | [[sklearn-集成学习|集成学习 →]]

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询