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

如何在R中按变量组构建Random Forest模型并获取组级RMSE

解决方案:按变量组计算RMSE指标

核心思路

要获取变量组对应的RMSE,我们可以通过置换目标变量组内所有变量的值,破坏该组与预测目标的关联,再用置换后的数据集进行预测并计算RMSE。这个RMSE与原模型基准RMSE的差值,能直接反映该变量组对预测性能的贡献;也可直接用置换后的RMSE衡量该组缺失时的预测误差。

修正原代码问题

  1. 原数据集生成逻辑错误:expand.grid生成504行数据,但后续变量仅提供9个值,会循环填充,先修正数据生成逻辑。
  2. 原配方语法错误:step_impute_median()调用方式不符合tidymodels规范,修正角色定义与预处理步骤。

完整实现代码

1. 数据准备与预处理

library(tidymodels)
library(parsnip)
library(ranger)
library(future)
library(dplyr)
library(purrr)

# 生成符合逻辑的数据集
set.seed(123)
countries <- expand.grid(country = c("Angola", "South Sudan", "Namibia"), 
                         year = 2006:2019, 
                         month = 1:12) %>%
  mutate(deaths = round(runif(nrow(.), 1000, 20000), 0),
         a = sample(c(6,7,4), nrow(.), replace = TRUE),
         b = sample(c(5,8,9), nrow(.), replace = TRUE),
         c = sample(c(2,20,80), nrow(.), replace = TRUE),
         d = sample(c(100,300,500), nrow(.), replace = TRUE))

# 划分训练/测试集
data_split <- initial_split(countries, prop = 0.7)
train_data <- training(data_split)
test_data  <- testing(data_split)

# 3折交叉验证方案
cv <- vfold_cv(train_data, v=3)

# 定义预处理配方,正确设置变量角色
my_recipe <- recipe(deaths ~ ., data = train_data) %>%
  add_role(country, year, month, new_role = "id vars") %>%
  add_role(a, b, new_role = "climate") %>%
  add_role(c, d, new_role = "conflict") %>%
  step_novel(all_nominal(), -has_role("id vars")) %>%
  step_impute_median(all_numeric(), -all_outcomes(), -has_role("id vars")) %>%
  step_dummy(all_nominal(), -has_role("id vars")) %>%
  prep()

2. 训练随机森林模型

# 定义RF模型
mod_rf <- rand_forest(trees = 1000) %>%
  set_engine("ranger", num.threads = parallel::detectCores(), importance = "permutation") %>%
  set_mode("regression")

# 构建工作流
wflow_rf <- workflow() %>%
  add_model(mod_rf) %>%
  add_recipe(my_recipe)

# 并行训练模型
plan(multisession)
fit_rf <- fit_resamples(
  wflow_rf,
  cv,
  metrics = metric_set(rmse, rsq),
  control = control_resamples(save_pred = TRUE, extract = function(x) extract_model(x))
)

3. 按变量组计算RMSE

编写函数对指定变量组进行置换,计算每个交叉验证折的RMSE后取平均值:

# 组置换RMSE计算函数
group_rmse <- function(fit_resamples, recipe, group_role) {
  # 获取目标组的变量名
  group_vars <- recipe$var_info %>%
    filter(role == group_role) %>%
    pull(variable) %>%
    as.character()
  
  # 遍历每个交叉验证折计算置换后RMSE
  map_dfr(fit_resamples$splits, function(split) {
    test_data <- assessment(split)
    # 置换组内所有变量
    permuted_test <- test_data %>%
      mutate(across(all_of(group_vars), ~sample(., replace = FALSE)))
    # 提取对应折的模型并预测
    model_idx <- which(fit_resamples$splits == split)
    model <- extract_model(fit_resamples$.extracts[[model_idx]][[1]])
    pred <- predict(model, new_data = bake(recipe, permuted_test))
    # 计算RMSE
    rmse_val <- rmse_vec(test_data$deaths, pred$.pred)
    tibble(group = group_role, rmse = rmse_val)
  }) %>%
    group_by(group) %>%
    summarise(mean_rmse = mean(rmse), .groups = "drop")
}

# 计算气候组、冲突组的RMSE
climate_rmse <- group_rmse(fit_rf, my_recipe, "climate")
conflict_rmse <- group_rmse(fit_rf, my_recipe, "conflict")

# 查看结果
bind_rows(climate_rmse, conflict_rmse)

4. 获取原模型基准RMSE

对比原模型的交叉验证RMSE,作为性能基准:

original_rmse <- fit_rf %>%
  collect_metrics() %>%
  filter(.metric == "rmse") %>%
  pull(mean)

cat("原模型基准RMSE:", original_rmse, "\n")

结果解释

  • 原模型RMSE是所有变量参与预测的基准误差。
  • 变量组置换后的RMSE若远大于基准值,说明该组对预测性能有显著贡献;差值越大,贡献越强。
  • 可通过两组RMSE的对比,判断气候变量组与暴力变量组对预测结果的影响程度。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 13:54:38