sklearn决策树高效获取各节点及叶节点对应数据记录的方法咨询
优化方案
1. 替换DataFrame传递为样本索引传递(核心优化,可降低70%+耗时)
现有方案的核心开销来自每次递归都生成新的DataFrame,pandas的布尔索引、新对象创建的开销远高于纯numpy操作。你完全可以只传递当前节点对应的样本索引数组,仅在需要某个节点的结构化数据时,再用索引从原DataFrame提取。
参考实现:
import numpy as np # 提前缓存全局属性,避免递归中重复取值 tree = clf.tree_ fn = [X.columns[i] if i != TREE_UNDEFINED else "undefined!" for i in tree.feature] # 提前转numpy数组做计算,列名映射成索引,跳过pandas列查询开销 X_arr = X.values col2idx = {col: i for i, col in enumerate(X.columns)} # 仅存储每个节点的样本索引,不存完整DataFrame node_indices = {} def recurse(node, sample_idx): # 原有节点判断逻辑可按需调整,需要DataFrame时再临时生成X.iloc[sample_idx] # if self.test_node(X.iloc[sample_idx]): # return node_indices[node] = sample_idx if tree.feature[node] != TREE_UNDEFINED: col_idx = col2idx[fn[node]] # 直接在numpy数组上做阈值判断,开销远低于pandas列操作 mask = X_arr[sample_idx, col_idx] <= tree.threshold[node] recurse(tree.children_left[node], sample_idx[mask]) recurse(tree.children_right[node], sample_idx[~mask]) # 初始传入全量样本索引 recurse(0, np.arange(X.shape[0]))
后续需要某节点的DataFrame时,再调用X.iloc[node_indices[node_id]]生成即可。
2. 延迟DataFrame生成
如果不是所有节点都需要输出DataFrame,不要在递归过程中做转换,仅在最终使用节点数据时再生成对应DataFrame,避免不必要的对象创建开销。
3. 可选项:递归改迭代
1万棵树的场景下,递归的函数调用开销会累计,可改用栈存(节点ID, 样本索引)元组的迭代方式遍历树,进一步降低调用开销。
补充说明:
关于你提到的「直接挂载训练时的节点样本」:scikit-learn的决策树默认仅存储每个节点的样本量、基尼系数等统计值,不会存储原始样本索引,从底层Cython实现中提取样本列表的拷贝成本反而高于上述索引传递方案。另外你测试中decision_path方案更慢是符合预期的,该方案时间复杂度为
O(样本数*节点数),而递归拆分的复杂度为O(样本数*树深),树深较小时前者开销要大得多。
内容的提问来源于stack exchange,提问作者mibm
相关产品推荐
相关产品推荐

