mhe*_*ler 4 python scikit-learn
如果对 max_depth、min_samples 等没有限制,有没有办法检索由 sklearn.tree.DecisionTreeClassifier 生成的最终节点数?
一旦你有了这棵树,你就可以访问它的内部tree_对象,以及这棵树的各种属性。源代码中描述的其中之一是node_count:
属性
- node_count : int
树中的节点数(内部节点 + 叶子)。
所以你可以这样做:
c = DecisionTreeClassifier(…)
c.fit(…)
n_nodes = c.tree_.node_count
Run Code Online (Sandbox Code Playgroud)
节点的其他各种属性存储在作为树对象属性的数组中,并由节点 id 索引。例如,它的value属性是一个节点分数数组,n_node_samples是每个节点的样本数数组。ericmjl的答案更详细地介绍了有关该表示的参考资料。您可以使用它来获取特定节点的值:
c = DecisionTreeClassifier(…)
c.fit(…)
value_i = c.tree_.value[i]
Run Code Online (Sandbox Code Playgroud)
部分原因是我自己没有决策树,所以我无法在这里提供具体的例子。不过,我认为深入研究源代码可能会有所帮助。
https://github.com/scikit-learn/scikit-learn/blob/master/sklearn/tree/_tree.pyx
sklearn的决策树有一个属性tree_,它是底层的Tree对象。Tree 对象在 _tree.pyx 类中定义。Tree 对象是:
“二叉决策树的基于数组的表示。二叉树表示为多个并行数组。
i每个数组的第 - 个元素保存有关节点的信息i。节点 0 是树的根。
回想一下,树中有一些内部节点,但决策树始终是二元的。因此,当您从节点 0 开始遍历时,您可以访问树属性children_left并children_right找出哪些节点是每个节点的子节点。该threshold属性可能对您也有用。
一些可能(强调可能)工作的伪代码是:
clf = DecisionTreeClassifier()
[... train and test...]
print(clf.tree_.node_count) #get the node count.
print(clf.tree_.children_left[node]) #where node is some integer
print(clf.tree_.children_right[node])
print(clf.tree_.threshold[node])
Run Code Online (Sandbox Code Playgroud)
希望这可以帮助。
| 归档时间: |
|
| 查看次数: |
5733 次 |
| 最近记录: |