如何基于测试数据集绘制已训练完成的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
相关产品推荐
相关产品推荐

