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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 23:45:00