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

如何从Tidymodels的Workflowset提取超参数并绘制RMSE、RSQ性能图?

问题描述

我希望基于Tidymodels中的workflowset绘制超参数的性能(RMSE和RSQ)图,但无法理清实现语法,想要复现指定图表,请问如何从我的race_results中提取超参数?以下是完整的Tidymodels建模及超参数调优的R代码示例:

# load the housing data and clean names
ames_data <- make_ames() %>%
  janitor::clean_names() %>% 
  mutate(sale_price_log = log10(sale_price))

# SPLIT INTO TRAINING AND TESTING DATA. STRATIFY BY SALE PRICE
ames_split <- rsample::initial_split(
  ames_data %>% select(-sale_price),
  prop = 0.8,
  strata = sale_price_log
)

# CREATE TRAINING AND TESTING OBJECTS FROM THE SPLIT OBJECT
ames_train <- training(ames_split)
names_test <- testing(ames_split)

# CREATE RESAMPLES TO CHOOSE AND COMPARE MODELS
set.seed(234)
names_folds <- vfold_cv(ames_train, strata = sale_price_log, v = 5)

# DEFINE PREPROCESSING RECIPES --------------------------------------------

base_rec <- recipe(sale_price_log ~ ., data = ames_train) %>%
  # APPLYING LOG TRANSFORMATION TO SALE_PRICE AND GR_LIV_AREA TO ADDRESS SKEWNESS
  step_log(gr_liv_area, base = 10) %>%
  # CREATE DUMMY VARIABLES FROM FACTOR COLUMNS
  step_dummy(all_nominal_predictors(), one_hot = TRUE)

normalise_rec <- recipe(sale_price_log ~ ., data = ames_train) %>%
  # REMOVE ANY COLUMNS WITH A SINGLE UNIQUE VALUE
  step_nzv(all_nominal_predictors()) %>%
  # HANDLING RARE FACTOR LEVELS IN NEIGHBORHOOD TO IMPROVE MODEL ROBUSTNESS
  step_other(all_nominal_predictors(), threshold = 0.05, other = "OTHER") %>%
  # STABILIZING VARIANCE AND NORMALIZING DISTRIBUTIONS FOR LOT_AREA AND GR_LIV_AREA
  step_YeoJohnson(all_numeric_predictors()) %>% 
  # NORMALIZING ALL NUMERIC PREDICTORS TO ENSURE THEY ARE ON A SIMILAR SCALE
  step_normalize(all_numeric_predictors()) %>%
  # CREATE DUMMY VARIABLES FROM FACTOR COLUMNS
  step_dummy(all_nominal_predictors(), one_hot = TRUE) %>%
  # REMOVE ANY COLUMNS WITH A SINGLE UNIQUE VALUE
  step_zv(all_predictors())

# PCA RECIPE
pca_rec <- recipe(sale_price_log ~ ., data = ames_train) %>% 
  # FOR UNSEEN FACTROR LEVELS, CREATE A NEW LEVEL CALLED "NEW"
  step_novel(all_nominal_predictors()) %>%
  # CREATE DUMMY VARIABLES FROM FACTOR COLUMNS
  step_dummy(all_nominal_predictors()) %>%
  # REMOVE ANY COLUMNS WITH A SINGLE UNIQUE VALUE
  step_zv(all_predictors()) %>%
  # NORMALIZING ALL NUMERIC PREDICTORS TO ENSURE THEY ARE ON A SIMILAR SCALE
  step_normalize(all_numeric_predictors()) %>%
  # CONVERT NUMERIC COLUMNS TO PRINCIPAL COMPONENTS
  step_pca(all_predictors(), threshold = 0.95)

# BUILD MODELS -----------------------------------------------------------

# DEFINE A BAGGED RANDOM FOREST MODEL
bagged_spec <- bag_tree(
  tree_depth = tune(),
  min_n = tune(),
  cost_complexity = tune()
) %>%
  set_mode("regression") %>%
  set_engine("rpart", times = 25L)

# DEFINE A RANGER RANDOM FOREST MODEL
rf_spec <-
  rand_forest(
    mtry = tune(),
    min_n = tune(),
    trees = 500
  ) %>%
  set_engine("ranger") %>%
  set_mode("regression")

# DEFINE AN XGBOOST MODEL
xgb_spec <- boost_tree(
  trees = 500,
  tree_depth = tune(),
  min_n = tune(),
  loss_reduction = tune(),
  sample_size = tune(),
  mtry = tune(),
  learn_rate = tune()
) %>%
  set_engine("xgboost", importance = TRUE) %>%
  set_mode("regression")

# DEFINE A BOOSTED TREE ENSEMBLE MODEL
bt_spec <-
  boost_tree(
    learn_rate = tune(),
    stop_iter = tune(),
    trees = 500
  ) %>%
  set_engine("lightgbm", num_leaves = tune()) %>%
  set_mode("regression")

# DEFINE A WORKFLOW SET ---------------------------------------------------

wflw_set <-
  workflow_set(
    preproc = list(base = base_rec, normalise = normalise_rec, pca = pca_rec),
    models = list(xgb = xgb_spec, bagged = bagged_spec, rf = rf_spec, bt = bt_spec),
    cross = TRUE
  )

# UPDATE MTRY PARAMETER FOR THE BASE XGBOOST
base_xgb_param <- wflw_set %>%
  extract_workflow(
    id = "base_xgb"
  ) %>%
  hardhat::extract_parameter_set_dials() %>%
  update(mtry = mtry(c(1, 308)))

base_rf_param <- wflw_set %>% 
  extract_workflow(
    id = "base_rf"
  ) %>%
  hardhat::extract_parameter_set_dials() %>%
  update(mtry = mtry(c(1, 308)))

# UPDATE MTRY PARAMETER FOR THE NORMALISED XGB MODEL
normalise_xgb_param <- wflw_set %>%
  extract_workflow(
    id = "normalise_xgb"
  ) %>%
  hardhat::extract_parameter_set_dials() %>%
  update(mtry = mtry(c(1, 284)))

# UPDATE MTRY PARAMETER FOR THE NORMALISED RF MODEL
normalise_rf_param <- wflw_set %>%
  extract_workflow(
    id = "normalise_rf"
  ) %>%
  hardhat::extract_parameter_set_dials() %>%
  update(mtry = mtry(c(1, 284)))

# UPDATE MTRY PARAMETER FOR THE PCA XGB MODEL
pca_xgb_param <- wflw_set %>%
  extract_workflow(
    id = "pca_xgb"
  ) %>%
  hardhat::extract_parameter_set_dials() %>%
  update(mtry = mtry(c(1, 5)))

# UPDATE MTRY PARAMETER FOR THE PCA XGB MODEL
pca_rf_param <- wflw_set %>%
  extract_workflow(
    id = "pca_rf"
  ) %>%
  hardhat::extract_parameter_set_dials() %>%
  update(mtry = mtry(c(1, 5)))

# UPDATE THE WORKFLOW SET WITH THE NEW PARAMETERS
wf_set_tune_list_finalize <- wflw_set %>%
  option_add(param_info = base_xgb_param, id = "base_xgb") %>%
  option_add(param_info = base_rf_param, id = "base_rf") %>%
  option_add(param_info = normalise_xgb_param, id = "normalise_xgb") %>%
  option_add(param_info = normalise_rf_param, id = "normalise_rf") %>%
  option_add(param_info = pca_xgb_param, id = "pca_xgb") %>%
  option_add(param_info = pca_rf_param, id = "pca_rf")

# SPECIFY THE TUNE GRID
race_ctrl <-
  control_race(
    save_pred = TRUE,
    parallel_over = "everything",
    save_workflow = TRUE
  )

# DETECT THE NUMBER OF CORES
cores <- parallel::detectCores(logical = FALSE)

# CREATE A SET OF COPIES OF R RUNNING IN PARALLEL AND COMMUNICATING VIA SOCKETS
cl <- makePSOCKcluster(cores)

# REGISTER THE PARALLEL BACKEND
doParallel::registerDoParallel(cores = cl)

# APPLY RACE ANOVA TUNING TO EACH WORKFLOW IN THE WORKFLOW SET
tictoc::tic()
race_results <- wf_set_tune_list_finalize %>%
  workflow_map(
    "tune_race_anova",
    seed = 123,
    resamples = ames_folds,
    grid = 5,
    control = race_ctrl,
    verbose = TRUE
  )
tictoc::toc()

# EXTRACT THE BEST RESULTS
best_results <- 
  race_results %>% 
  extract_workflow_set_result("base_xgb") %>% 
  select_best(metric = "rmse")
解决方案

一、从race_results提取超参数与性能指标

race_results是workflow_set对象,每个工作流的调优结果存储在result列中,可通过以下方式提取超参数和对应性能数据:

1. 提取单个工作流的数据

以base_xgb为例,先提取调优结果,再整合超参数和性能指标:

library(dplyr)
library(tune)

# 提取base_xgb的调优结果对象
base_xgb_res <- race_results %>% extract_workflow_set_result("base_xgb")

# 整合超参数与RMSE、RSQ等性能指标
base_xgb_tidy <- base_xgb_res %>%
  collect_metrics() %>%
  left_join(collect_parameters(base_xgb_res), by = ".config")

base_xgb_tidy包含了该工作流所有超参数组合的性能均值、标准差,以及对应的超参数取值。

2. 批量提取所有工作流的数据

如果需要处理所有工作流,用purrr批量整合:

library(purrr)
library(tidyr)

all_tuned_data <- race_results %>%
  mutate(
    tidy_data = map(result, function(res) {
      collect_metrics(res) %>%
        left_join(collect_parameters(res), by = ".config") %>%
        mutate(workflow_id = res$id[1])
    })
  ) %>%
  select(workflow_id, tidy_data) %>%
  unnest(tidy_data)

all_tuned_data包含所有工作流的超参数、性能指标及工作流ID,适合跨工作流对比分析。

二、绘制超参数性能图

用ggplot2绘制超参数与性能的关系,以下是几种常见场景的实现:

1. 单个超参数对性能的影响

比如查看XGBoost的learn_rate对RMSE的影响:

library(ggplot2)

base_xgb_tidy %>%
  filter(.metric == "rmse") %>%
  ggplot(aes(x = learn_rate, y = mean)) +
  geom_point(size = 2, color = "#2E86AB") +
  geom_errorbar(aes(ymin = mean - std_err, ymax = mean + std_err), width = 0.001) +
  labs(
    title = "XGBoost学习率对RMSE的影响",
    x = "学习率(learn_rate)",
    y = "RMSE均值(含标准误)"
  ) +
  theme_minimal()

2. 跨工作流对比同一超参数的性能

对比不同预处理下XGBoost的mtry对RMSE的影响:

all_tuned_data %>%
  filter(.metric == "rmse", workflow_id %in% c("base_xgb", "normalise_xgb", "pca_xgb")) %>%
  ggplot(aes(x = mtry, y = mean, color = workflow_id)) +
  geom_point(size = 2) +
  geom_errorbar(aes(ymin = mean - std_err, ymax = mean + std_err), width = 0.1) +
  labs(
    title = "不同预处理下XGBoost的mtry对RMSE的影响",
    x = "mtry",
    y = "RMSE均值",
    color = "工作流"
  ) +
  theme_minimal() +
  facet_wrap(~workflow_id, scales = "free_x")

3. 两个超参数组合的热力图

查看XGBoost的tree_depth和min_n对RMSE的联合影响:

base_xgb_tidy %>%
  filter(.metric == "rmse") %>%
  ggplot(aes(x = tree_depth, y = min_n, fill = mean)) +
  geom_tile(color = "white") +
  geom_text(aes(label = round(mean, 4)), color = "white", size = 3) +
  scale_fill_viridis_c(option = "mako") +
  labs(
    title = "XGBoost树深与最小样本量对RMSE的影响",
    x = "树深(tree_depth)",
    y = "最小样本量(min_n)",
    fill = "RMSE均值"
  ) +
  theme_minimal()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 15:24:52