R语言如何基于rpart包结果绘制二维协变量空间递归分割图
R语言绘制rpart树二维递归分割分区图
实现方案
不需要自定义rpart_splits函数做额外转换,直接读取rpart对象内置的节点与分割信息,通过递归遍历树结构追踪每个分割对应的区域边界,即可绘制出《统计学习基础》中展示的轴对齐二元分割效果图,函数可直接传入rpart建模结果调用。
完整函数代码
library(rpart) plot_rpart_2d_partition <- function( tree, x_var = "X1", y_var = "X2", point_col = rgb(0, 0, 0, 0.3), split_col = "black", split_lwd = 1.5, show_leaf_pred = TRUE, pred_col = "red", pred_cex = 0.8 ) { # 提取建模数据与变量范围 model_data <- eval(tree$call$data, envir = parent.frame()) xlim <- range(model_data[[x_var]]) ylim <- range(model_data[[y_var]]) # 提前提取所有非叶节点的主分割信息 frame <- tree$frame is_leaf <- frame$var == "<leaf>" split_cumpos <- cumsum(c(1, frame$ncompete + frame$nsurrogate + !is_leaf)) main_split_list <- list() for (i in seq_along(is_leaf)) { if (!is_leaf[i]) { node_id <- as.character(rownames(frame)[i]) split_row <- split_cumpos[i] main_split_list[[node_id]] <- list( split_var = as.character(frame$var[i]), split_threshold = tree$splits[split_row, "index"] ) } } # 绘制底图散点 plot( model_data[[x_var]], model_data[[y_var]], pch = 16, col = point_col, xlim = xlim, ylim = ylim, xlab = x_var, ylab = y_var, main = "CART递归二元分割二维分区" ) # 递归遍历节点绘制分割线 traverse_draw <- function(node_id, x_left, x_right, y_bottom, y_top) { node_name <- as.character(node_id) # 到达叶节点,按需标注预测值 if (is_leaf[node_name]) { if (show_leaf_pred) { pred_val <- round(frame[node_name, "yval"], 2) text( mean(c(x_left, x_right)), mean(c(y_bottom, y_top)), labels = pred_val, cex = pred_cex, col = pred_col ) } return() } # 读取当前节点分割信息 current_split <- main_split_list[[node_name]] split_v <- current_split$split_var thresh <- current_split$split_threshold # 按分割变量绘制对应方向的分割线,再递归处理左右子节点 if (split_v == x_var) { # 垂直分割(沿X轴切分) segments( x0 = thresh, y0 = y_bottom, x1 = thresh, y1 = y_top, col = split_col, lwd = split_lwd ) traverse_draw(2 * node_id, x_left, thresh, y_bottom, y_top) traverse_draw(2 * node_id + 1, thresh, x_right, y_bottom, y_top) } else if (split_v == y_var) { # 水平分割(沿Y轴切分) segments( x0 = x_left, y0 = thresh, x1 = x_right, y1 = thresh, col = split_col, lwd = split_lwd ) traverse_draw(2 * node_id, x_left, x_right, y_bottom, thresh) traverse_draw(2 * node_id + 1, x_left, x_right, thresh, y_top) } } # 从根节点(ID=1)开始遍历 traverse_draw(1, xlim[1], xlim[2], ylim[1], ylim[2]) }
调用测试
直接使用提供的示例代码即可生成目标图像:
# 生成模拟数据 set.seed(1975) n <- 5000 p <- 2 X <- matrix(sample(seq(0, 1, by = 0.01), n * p, replace = TRUE), ncol = p) Y <- X[, 1] + 2 * X[, 2] + rnorm(n) # 训练CART树 tree <- rpart( Y ~ ., data = data.frame(Y, X), method = "anova", control = rpart.control(cp = 0, maxdepth = 2) ) # 绘制分区图 plot_rpart_2d_partition(tree)
效果说明
- 生成的图像和《统计学习基础》第306页示例格式一致:底层是二维协变量的散点分布,黑色实线为轴对齐分割线,每个叶节点区域中心标注该区域的预测值
- 分割位置和手动提取的结果完全匹配:根节点在X2=0.565处水平切分,左子节点在X2=0.265处二次水平切分,右子节点在X1=0.645处垂直切分
- 若不需要叶节点预测值标注,调用时设置参数
show_leaf_pred = FALSE即可;若协变量名称不是默认的X1、X2,传入对应列名给x_var、y_var参数即可适配
内容的提问来源于stack exchange,提问作者riccardo-df
相关产品推荐
相关产品推荐

