如何用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
相关产品推荐
相关产品推荐

