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

Spark MLlib(Scala)决策树能否返回样本所属叶节点索引?

在Spark MLlib中获取样本对应的叶节点索引

当然有!Spark MLlib提供了和Scikit-learn apply方法功能一致的实现,不过要分你使用的是**新的ML(DataFrame API)还是旧的mllib(RDD API)**两种情况,具体如下:

1. 使用Spark ML(DataFrame API,推荐)

如果你用的是Spark 2.3及以上版本,DecisionTreeClassificationModel和DecisionTreeRegressionModel都内置了predictLeaf()方法,直接就能返回每个样本被预测到的叶节点索引。

举个Scala代码示例:

import org.apache.spark.ml.classification.DecisionTreeClassificationModel
import org.apache.spark.sql.DataFrame

// 假设你已经训练好决策树分类模型
val trainedDtModel: DecisionTreeClassificationModel = ...
// 测试数据集
val testData: DataFrame = ...

// 获取每个样本对应的叶节点索引,返回一个Column
val leafIndicesCol = trainedDtModel.predictLeaf(testData("features"))
// 可以把这个列添加到原DataFrame中查看
val resultData = testData.withColumn("leaf_node_index", leafIndicesCol)

predictLeaf()返回的是一个整数类型的列,每个值就是对应样本最终落到的叶节点的内部索引,和Scikit-learn的apply输出逻辑完全一致。

2. 使用旧的MLlib(RDD API)

如果你的代码基于旧的RDD风格的mllib API,DecisionTreeModel(分类或回归)提供了predictNode()方法,它会返回样本所在节点的ID——对于叶节点来说,这个ID就是对应的叶节点索引。

示例代码:

import org.apache.spark.mllib.tree.model.DecisionTreeModel
import org.apache.spark.mllib.linalg.Vector
import org.apache.spark.rdd.RDD

// 训练好的RDD风格决策树模型
val trainedDtModel: DecisionTreeModel = ...
// 测试RDD,每个元素是特征向量
val testRDD: RDD[Vector] = ...

// 映射得到每个样本的叶节点索引
val leafNodeIndicesRDD = testRDD.map { features =>
  trainedDtModel.predictNode(features)
}

需要注意的是,这个方法返回的节点ID包含内部非叶节点的编号,但当样本被分到叶节点时,返回的就是叶节点的索引,和你需要的效果一致。

补充说明

  • 叶节点索引是模型训练时内部生成的唯一标识,不同模型的索引不具备可比性;
  • 如果你使用的是Spark 2.3以下的ML API,可能需要通过解析模型的toDebugString来手动实现路径追踪,但这种方式比较繁琐,建议尽量升级到较新版本的Spark来使用内置方法。

内容的提问来源于stack exchange,提问作者Yifan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 02:29:06