使用rpart拟合模型后,如何获取验证/测试数据集的where向量?
获取rpart模型验证/测试集的叶节点编号(类似where向量)
当然可以实现!rpart模型自带的where向量确实只对应训练数据集的叶节点,但我们可以用predict()函数轻松拿到验证集或测试集样本对应的叶节点编号,效果和训练集的where完全一致。
具体步骤:
- 先拟合你的rpart模型
library(rpart) # 用鸢尾花数据集举例子,替换成你的数据即可 model <- rpart(Species ~ ., data = iris_train)
训练集的叶节点编号可以直接用model$where查看。
- 对验证/测试集使用
predict(),指定type="where"参数
# 假设你已经有测试集iris_test test_where <- predict(model, newdata = iris_test, type = "where")
这里的test_where就是测试集每条样本对应的叶节点编号,和训练集model$where的编号体系完全统一——也就是说,如果某条测试样本和某条训练样本落在同一个叶节点,它们的编号是相同的。
完整示例:
# 1. 拆分训练集和测试集 set.seed(123) # 设置随机种子保证可复现 train_idx <- sample(nrow(iris), 100) iris_train <- iris[train_idx, ] iris_test <- iris[-train_idx, ] # 2. 拟合rpart模型 model <- rpart(Species ~ ., data = iris_train) # 3. 获取训练集叶节点编号 train_where <- model$where head(train_where) # 查看前6个训练样本的叶节点 # 4. 获取测试集叶节点编号 test_where <- predict(model, newdata = iris_test, type = "where") head(test_where) # 查看前6个测试样本的叶节点
这个方法是rpart官方支持的,完全可以满足你获取新数据集对应叶节点的需求~
内容的提问来源于stack exchange,提问作者Meng zhao
相关产品推荐
相关产品推荐

