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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 18:45:38