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

如何基于测试数据集绘制已训练完成的CART决策树?

基于测试集复用CART分裂规则并更新节点统计量绘图

要实现保留训练好的CART模型分裂规则,仅用测试集更新节点决策与占比的需求,我们可以通过复用原模型结构、替换节点统计信息的方式实现,具体步骤如下:

核心思路

rpart模型的分裂规则(如节点分裂变量、阈值)存储在模型对象的split等字段中,节点的样本统计量(类别计数、占比、预测类别)则存在frame字段。我们只需保留分裂规则,用测试集重新计算统计量并替换,即可得到基于测试集的模型可视化结果。

完整实现代码

library(plyr) 
library(dplyr) 
library(rpart) 
library(rpart.plot) 

# 数据预处理(复用原有代码)
stagec <- mutate(stagec, progstat = factor(pgstat, levels = 0:1, labels = c("No", "Prog")))
set.seed(1234)
train = sample(1:nrow(stagec), 73)
data_train <- stagec[train,]
data_test <- stagec[-train,]

# 训练CART模型(复用原有代码)
cfit <- rpart(progstat ~ age + eet + g2 + grade + gleason + ploidy, data = data_train, method = 'class')

# ------------------------------
# 关键步骤:用测试集更新模型节点统计量
# ------------------------------
# 1. 获取测试集样本在原模型中的节点归属
test_nodes <- predict(cfit, data_test, type = "matrix")[,1]

# 2. 计算每个节点的测试集统计量
node_stats <- data_test %>%
  mutate(node = test_nodes) %>%
  group_by(node) %>%
  summarise(
    n_total = n(),
    count_no = sum(progstat == "No"),
    count_prog = sum(progstat == "Prog"),
    .groups = "drop"
  ) %>%
  mutate(
    pred_class = ifelse(count_no > count_prog, "No", "Prog"),
    pred_code = match(pred_class, levels(data_test$progstat))  # 转换为rpart内部编码
  )

# 3. 复制原模型并替换节点统计信息
cfit_test <- cfit
for (row in 1:nrow(node_stats)) {
  node_id <- as.character(node_stats$node[row])
  frame_idx <- which(rownames(cfit_test$frame) == node_id)
  
  # 更新节点样本总数
  cfit_test$frame$n[frame_idx] <- node_stats$n_total[row]
  # 更新各类别计数(yval2第2列对应"No",第3列对应"Prog")
  cfit_test$frame$yval2[frame_idx, 2] <- node_stats$count_no[row]
  cfit_test$frame$yval2[frame_idx, 3] <- node_stats$count_prog[row]
  # 更新预测类别编码
  cfit_test$frame$yval[frame_idx] <- node_stats$pred_code[row]
}

# 4. 绘制基于测试集的决策树
rpart.plot(cfit_test, extra=104)

关键细节说明

  • predict(cfit, data_test, type = "matrix")返回的矩阵第一列是每个测试样本对应的节点编号,通过这个可以将测试样本映射到原模型的节点中。
  • cfit$frame是rpart存储节点信息的核心数据框,我们重点更新n(节点样本数)、yval2(类别计数矩阵)、yval(预测类别编码)三个字段,其余字段保留原模型的分裂规则。
  • extra=104参数会显示节点的样本占比和预测类别,绘图时会自动使用更新后的测试集统计量。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 01:27:10