【发布时间】:2017-03-26 07:25:00
【问题描述】:
我在 Spark 中有一个 RandomForestClassifierModel。使用 .toDebugString() 输出以下内容
Tree 0 (weight 1.0):
If (feature 0 in {1.0,2.0,3.0})
If (feature 3 in {2.0,3.0})
If (feature 8 <= 55.3)
.
.
Else (feature 0 not in {1.0,2.0,3.0})
.
.
Tree 1 (weight 1.0):
.
.
...etc
我想查看通过模型的实际数据,例如
Tree 0 (weight 1.0):
If (feature 0 in {1.0,2.0,3.0}) 60%
If (feature 3 in {2.0,3.0}) 57%
If (feature 8 <= 55.3) 22%
.
.
Else (feature 0 not in {1.0,2.0,3.0}) 40%
.
.
Tree 1 (weight 1.0):
.
...etc
通过查看每个节点中标签的概率,我可以看到数据(数千条记录)在树中最有可能遵循哪些路径,这将是非常好的洞察力!
我在这里找到了一个很棒的答案:Spark MLib Decision Trees: Probability of labels by features?
不幸的是,答案中的方法使用了 MLlib API,经过大量尝试,我未能使用 DataFrame API 复制它,它具有类 Node 和 Split 的不同实现:(
【问题讨论】:
-
也许您可以通过查看 ml 包的代码来尝试调整原始答案:github.com/apache/spark/blob/branch-2.0/mllib/src/main/scala/…
-
我查看了代码。我不确定这是否可能(至少与其他答案的策略不同),因为每个节点中的拆分比率位于“impurityStats”属性中,该属性是 ml 包的私有属性。也许可以使用 Node 的可见属性通过 ImpurityCalculator 创建此属性,但我找不到方法。
-
@DanieldePaula 感谢您对此进行调查。我宁愿不重构我的整个管道以使用 mllib。我会尝试使用您的最后建议找到一种方法。到目前为止,我可以将森林中的每一棵树都放在一个数组中。我想使用 API 来做到这一点,而不必重新编写大量的类。如果您碰巧想出任何其他解决方案,请告诉我!
-
再看一点,我发现它有一张未完成的票:issues.apache.org/jira/browse/SPARK-3727,但还没有完成。所以,显然,在 ml 包中还没有对它的支持。
标签: scala apache-spark apache-spark-sql spark-dataframe apache-spark-mllib