You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何从DecisionTreeRegressor中提取深度为1的左子树样本?

如何提取决策回归树指定子树对应的样本?

scikit-learn的DecisionTreeRegressor确实没有直接提取指定子树样本的方法,但可以通过decision_path()方法实现这个需求,核心思路是通过样本的决策路径筛选出到达目标节点的样本。

具体步骤

  1. 确定目标节点的编号:决策树的节点按广度优先规则编号,根节点为0,根节点的第一个左子节点(深度1)编号为1,右子节点为2,以此类推。如果不确定节点编号,可以用可视化工具查看节点ID。
  2. 获取决策路径矩阵:decision_path()会返回一个稀疏矩阵,每行代表一个样本,每列代表一个节点,矩阵中的非零值表示该样本在预测时经过了对应节点。
  3. 筛选目标样本:从矩阵中筛选出经过目标节点的样本索引,再用索引提取对应的样本数据。

代码示例

from sklearn.tree import DecisionTreeRegressor
from sklearn.datasets import load_boston
import numpy as np
from sklearn.tree import plot_tree
import matplotlib.pyplot as plt

# 1. 训练决策回归树
X, y = load_boston(return_X_y=True)
reg = DecisionTreeRegressor(max_depth=3, random_state=42)
reg.fit(X, y)

# 可选:可视化树结构,确认目标节点编号(node_ids=True显示节点ID)
plt.figure(figsize=(12, 8))
plot_tree(reg, filled=True, node_ids=True, feature_names=load_boston().feature_names)
plt.show()

# 2. 获取所有样本的决策路径
decision_paths = reg.decision_path(X)

# 3. 指定目标节点:深度1的第一个左子节点,编号为1
target_node_id = 1

# 筛选出经过目标节点的样本索引
sample_mask = decision_paths[:, target_node_id].toarray().flatten() > 0
target_sample_indices = np.where(sample_mask)[0]

# 4. 提取对应的样本和标签
target_samples = X[target_sample_indices]
target_labels = y[target_sample_indices]

print(f"目标子树对应的样本数量:{len(target_samples)}")

补充说明

  • 如果树结构更复杂,可通过reg.tree_.children_left和reg.tree_.children_right属性遍历节点关系,动态定位目标节点。比如reg.tree_.children_left[0]就是根节点的左子节点编号,和我们指定的1一致。
  • 稀疏矩阵操作时,用toarray()转为密集矩阵后筛选更直观,若样本量极大,推荐用稀疏矩阵的原生方法(如getnnz())提升效率。

内容的提问来源于stack exchange,提问作者SSSOF

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.04 01:43:34