L8058与L8168校色文件为何不能混用?硬件级色彩原理揭秘
2026/9/15 5:22:31
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)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()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 ...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')决策树容易过拟合,通过以下参数控制:
# 限制树的生长model=DecisionTreeClassifier(max_depth=5,# 限制深度min_samples_split=20,# 分裂所需最少样本min_samples_leaf=10,# 叶节点最少样本max_leaf_nodes=50,# 限制叶节点数量min_impurity_decrease=0.01,# 不纯度降低阈值)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-集成学习|集成学习 →]]