【问题标题】:How to view Random Forest statistics in Spark (scala)如何在 Spark (scala) 中查看随机森林统计信息
【发布时间】: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


【解决方案1】:

我昨天发现有用的一种方法是我可以使用 spark.read.parquet() 函数从模型/数据文件中读取输出。这样,关于某个节点的所有信息都可以作为一个完整的数据帧来检索。

`val modelPath = "some/path/to/your/model"
val dataPath = modelPath + "/data"    
val nodeData: DataFrame = spark.read.parquet(dataPath)
nodeData.show(500,false)
nodeData.printSchema()`

然后你可以用信息重建树。希望对您有所帮助。

【讨论】:

    猜你喜欢
    • 2017-06-13
    • 1970-01-01
    • 1970-01-01
    • 2012-07-20
    • 2020-11-30
    • 2016-01-28
    • 2018-06-06
    • 2018-12-31
    • 2017-12-04
    相关资源
    最近更新 更多