如何在R中绘制RandomForest的决策边界图并基于ggplot自定义
用ggplot绘制RandomForest三类分类的决策边界图(基于主成分)
核心思路
因为有3个主成分(PC),我们可以通过两两组合生成散点图+决策边界的方式完整展示分类效果,具体步骤为:
- 训练RandomForest分类模型
- 针对每一对PC,生成覆盖变量取值范围的网格数据
- 用训练好的RF模型预测网格数据的分类结果
- 用ggplot将原始数据散点与网格预测的边界填充/线条结合展示
完整代码示例
假设你的数据集名为df,其中PC1/PC2/PC3是主成分变量,class是三类目标分类变量。
1. 训练RandomForest模型
library(randomForest) library(ggplot2) library(dplyr) library(tidyr) # 训练RF模型,树数量可按需调整 rf_model <- randomForest(class ~ PC1 + PC2 + PC3, data = df, ntree = 500)
2. 编写边界生成工具函数
这个函数可以快速处理任意一对PC的网格生成与预测:
generate_decision_boundary <- function(pc_x, pc_y, model, data) { # 创建当前PC对的取值网格 x_range <- seq(min(data[[pc_x]]), max(data[[pc_x]]), length.out = 100) y_range <- seq(min(data[[pc_y]]), max(data[[pc_y]]), length.out = 100) # 构造包含三个PC的网格数据,固定第三个PC为中位数(可换成均值/分位数) grid <- expand.grid( PC1 = if(pc_x == "PC1") x_range else if(pc_y == "PC1") y_range else data$PC1, PC2 = if(pc_x == "PC2") x_range else if(pc_y == "PC2") y_range else data$PC2, PC3 = if(pc_x == "PC3") x_range else if(pc_y == "PC3") y_range else data$PC3 ) fixed_pc <- setdiff(c("PC1","PC2","PC3"), c(pc_x, pc_y)) grid[[fixed_pc]] <- median(data[[fixed_pc]]) # 预测网格的分类结果 grid$pred_class <- predict(model, newdata = grid) return(grid) }
3. 绘制三组PC对的边界图
# 生成三对PC的边界数据 boundary_pc1_pc2 <- generate_decision_boundary("PC1", "PC2", rf_model, df) boundary_pc1_pc3 <- generate_decision_boundary("PC1", "PC3", rf_model, df) boundary_pc2_pc3 <- generate_decision_boundary("PC2", "PC3", rf_model, df) # PC1 vs PC2 决策边界图 plot_pc1_pc2 <- ggplot() + geom_tile(data = boundary_pc1_pc2, aes(x = PC1, y = PC2, fill = pred_class), alpha = 0.3) + geom_point(data = df, aes(x = PC1, y = PC2, color = class), size = 2) + labs(title = "PC1 vs PC2 决策边界", x = "PC1", y = "PC2") + scale_fill_brewer(palette = "Set2") + scale_color_brewer(palette = "Set2") + theme_minimal() # PC1 vs PC3 决策边界图 plot_pc1_pc3 <- ggplot() + geom_tile(data = boundary_pc1_pc3, aes(x = PC1, y = PC3, fill = pred_class), alpha = 0.3) + geom_point(data = df, aes(x = PC1, y = PC3, color = class), size = 2) + labs(title = "PC1 vs PC3 决策边界", x = "PC1", y = "PC3") + scale_fill_brewer(palette = "Set2") + scale_color_brewer(palette = "Set2") + theme_minimal() # PC2 vs PC3 决策边界图 plot_pc2_pc3 <- ggplot() + geom_tile(data = boundary_pc2_pc3, aes(x = PC2, y = PC3, fill = pred_class), alpha = 0.3) + geom_point(data = df, aes(x = PC2, y = PC3, color = class), size = 2) + labs(title = "PC2 vs PC3 决策边界", x = "PC2", y = "PC3") + scale_fill_brewer(palette = "Set2") + scale_color_brewer(palette = "Set2") + theme_minimal() # 组合三个图(需加载gridExtra包) library(gridExtra) grid.arrange(plot_pc1_pc2, plot_pc1_pc3, plot_pc2_pc3, ncol = 2)
自定义调整技巧
- 边界平滑度:修改
generate_decision_boundary函数中length.out的数值,数值越大网格越密,边界越平滑,但计算耗时会增加。 - 第三个PC的取值:可以将中位数替换为均值、分位数,甚至制作动画展示不同取值下的边界变化。
- 样式优化:
- 把
geom_tile换成geom_contour可只绘制边界线,不做区域填充 - 调整
alpha参数控制填充区域的透明度 - 用
scale_fill_manual/scale_color_manual自定义分类颜色
- 把
- 概率可视化:如果需要展示分类置信度,可使用
predict(rf_model, newdata = grid, type = "prob")获取各类别概率,再用geom_tile填充概率渐变效果。
内容的提问来源于stack exchange,提问作者Myriad
相关产品推荐
相关产品推荐

