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

如何用R中tidymodels调优模型预测测试集置信区间?报错解决

解决tidymodels中last_fit报错及测试集置信区间预测问题

错误原因

last_fit()函数要求传入rsplit对象(即initial_split()生成的拆分对象),但你传入的是训练集数据框Sac_train,因此触发Each element of splits must be an rsplit object错误。last_fit()的作用是用整个训练集拟合最终模型,同时自动在测试集上完成评估,必须依赖初始拆分对象来区分训练/测试集。

修正步骤及完整代码

1. 修正last_fit调用

将last_fit(Sac_train)替换为last_fit(data_split),传入初始拆分的rsplit对象:

# Data splitting
data(Sacramento, package = "modeldata")
set.seed(123)
data_split <- initial_split(Sacramento, prop = 0.75, strata = price)
Sac_train <- training(data_split)
Sac_test <- testing(data_split)

# Build the model
rf_mod <- rand_forest(mtry = tune(), min_n = tune(), trees = 1000) %>% 
          set_engine("ranger", importance = "permutation") %>% 
          set_mode("regression")

# Create the recipe
Sac_recipe <- recipe(price ~ ., data = Sac_train) %>% 
              step_rm(zip, latitude, longitude) %>% 
              step_corr(all_numeric_predictors(), threshold = 0.85) %>% 
              step_zv(all_numeric_predictors()) %>% 
              step_normalize(all_numeric_predictors()) %>%
              step_dummy(all_nominal_predictors())

# Create the workflow
rf_workflow <- workflow() %>% 
               add_model(rf_mod) %>% 
               add_recipe(Sac_recipe)

# Train and Tune the model
set.seed(123)
Sac_folds <- vfold_cv(Sac_train, v = 10, repeats = 2, strata = price)

rf_res <- rf_workflow %>% 
          tune_grid(grid = 2*2,
                    resamples = Sac_folds, 
                    control = control_grid(save_pred = TRUE),
                    metrics = metric_set(rmse))

# Extract the best model
rf_best <- rf_res %>%
           select_best(metric = "rmse")

# Last fit - 修正:传入rsplit对象data_split
last_rf_workflow <- rf_workflow %>% 
                    finalize_workflow(rf_best)

last_rf_fit <- last_rf_workflow %>% 
               last_fit(data_split)  # 这里是关键修正

2. 获取测试集的置信区间预测

有两种方式可以得到测试集的置信区间:

方式一:从last_fit结果中直接提取

last_fit()会自动在测试集上执行预测,只需指定type = "conf_int"即可提取置信区间:

# 提取测试集的预测结果(包含真实值、预测值、置信区间上下限)
test_predictions <- last_rf_fit %>% 
  collect_predictions(type = "conf_int")

# 查看结果
head(test_predictions)

方式二:手动对测试集进行预测

如果需要单独对Sac_test执行预测,先提取拟合好的工作流,再调用predict():

# 提取已拟合的完整工作流
fitted_workflow <- last_rf_fit %>% extract_workflow()

# 对测试集预测置信区间
test_conf_int <- predict(fitted_workflow, Sac_test, type = "conf_int")

# 合并测试集原始数据与置信区间
Sac_test_with_ci <- bind_cols(Sac_test, test_conf_int)

注意事项

  • 使用ranger引擎的随机森林支持回归任务的置信区间预测,无需额外配置引擎参数,只需在predict()时指定type = "conf_int"。
  • collect_predictions()返回的结果包含.pred(预测值)、.pred_lower(置信区间下限)、.pred_upper(置信区间上限)以及原始测试集的响应变量price。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 22:40:40