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

在Tidymodels的workflow_set中使用tune_bayes调参报错求助

问题解决:tidymodels workflow_set 中贝叶斯调参的 mtry 参数错误

错误原因

报错提示mtry参数存在未知值,是因为rand_forest和boost_tree的mtry参数默认没有设定范围上限(上限取决于训练集的预测变量数量),而workflow_map调用tune_bayes时不会自动根据训练数据确定这个范围,必须手动通过finalize()明确参数范围。

解决方案

方案一:在模型spec中直接finalize mtry参数(推荐)

先预处理recipe得到最终的训练数据结构,再在定义模型时直接为mtry参数设定finalized的范围,这样workflow_map可以直接调用tune_bayes无需额外配置。

完整修改代码:

library(tidymodels)
library(tidyverse)
library(stacks) 

data("tree_frogs")

tree_frogs <- tree_frogs %>% select(-c(clutch, latency)) 
set.seed(1)

tree_frogs_split <- initial_split(tree_frogs)
tree_frogs_train <- training(tree_frogs_split)
tree_frogs_test  <- testing(tree_frogs_split)

tree_frogs_rec <- 
  recipe(reflex ~ ., data = tree_frogs_train) %>%
  step_dummy(all_nominal(), -reflex) %>%
  step_zv(all_predictors())

# 预处理recipe,得到 bake 后的训练数据,用于确定mtry的上限
tree_frogs_rec_prepped <- prep(tree_frogs_rec)
tree_frogs_train_baked <- bake(tree_frogs_rec_prepped, tree_frogs_train)

set.seed(1)
folds <- rsample::vfold_cv(tree_frogs_train, v = 5)

## 多分类模型
mlt_spec <- 
  multinom_reg(penalty = tune(), mixture = 1) %>% 
  set_engine("glmnet") %>% 
  set_mode("classification")

## 随机森林 - 已finalize mtry参数
rand_forest_spec <- 
  rand_forest(
    mtry = tune(dials::mtry() %>% finalize(tree_frogs_train_baked)),
    min_n = tune(),
    trees = 500
  ) %>%
  set_mode("classification") %>%
  set_engine("ranger")

## XGBoost - 已finalize mtry参数
xgb_spec <- 
  boost_tree(
    trees = 1000,
    min_n = tune(),
    learn_rate = tune(),
    loss_reduction = tune(),
    sample_size = tune(),
    mtry = tune(dials::mtry() %>% finalize(tree_frogs_train_baked)),
    tree_depth = tune()
  ) %>%
  set_engine("xgboost") %>%
  set_mode("classification")

all_workflows <- 
  workflow_set(
    preproc = list("basic_rec" = tree_frogs_rec),
    models = list(mlt = mlt_spec, rf = rand_forest_spec, xgb = xgb_spec)
  )

ctrl_bayes = control_bayes(save_pred = TRUE,
                           parallel_over = "everything",
                           save_workflow = TRUE)

bayes_results <- all_workflows %>%
  workflow_map(
    seed = 1,
    fn = "tune_bayes",
    resamples = folds,
    initial = 10,
    control = ctrl_bayes,
    metrics = metric_set(pr_auc, roc_auc, accuracy) # 匹配你需要的评估指标
  )

bayes_results %>% 
  rank_results() %>% 
  filter(.metric == "pr_auc") %>% 
  select(model, .config, pr_auc = mean, rank)

方案二:拆分workflow_set,分别传递finalized的param_info

如果不想修改模型spec,可以先为每个模型提取并finalize参数集,拆分workflow_set后单独训练,最后合并结果:

# 提取并finalize各模型的参数集
mlt_params <- extract_parameter_set_dials(mlt_spec)
rf_params <- extract_parameter_set_dials(rand_forest_spec) %>% finalize(tree_frogs_train_baked)
xgb_params <- extract_parameter_set_dials(xgb_spec) %>% finalize(tree_frogs_train_baked)

# 分别训练每个模型
mlt_res <- all_workflows %>%
  filter(wflow_id == "basic_rec_mlt") %>%
  pull(workflow) %>%
  first() %>%
  tune_bayes(
    resamples = folds,
    param_info = mlt_params,
    initial = 10,
    control = ctrl_bayes,
    metrics = metric_set(pr_auc, roc_auc, accuracy)
  )

rf_res <- all_workflows %>%
  filter(wflow_id == "basic_rec_rf") %>%
  pull(workflow) %>%
  first() %>%
  tune_bayes(
    resamples = folds,
    param_info = rf_params,
    initial = 10,
    control = ctrl_bayes,
    metrics = metric_set(pr_auc, roc_auc, accuracy)
  )

xgb_res <- all_workflows %>%
  filter(wflow_id == "basic_rec_xgb") %>%
  pull(workflow) %>%
  first() %>%
  tune_bayes(
    resamples = folds,
    param_info = xgb_params,
    initial = 10,
    control = ctrl_bayes,
    metrics = metric_set(pr_auc, roc_auc, accuracy)
  )

# 合并结果为workflow_set格式
bayes_results <- bind_rows(
  mlt_res %>% mutate(wflow_id = "basic_rec_mlt"),
  rf_res %>% mutate(wflow_id = "basic_rec_rf"),
  xgb_res %>% mutate(wflow_id = "basic_rec_xgb")
) %>%
  as_workflow_set_result()

# 后续分析
bayes_results %>% 
  rank_results() %>% 
  filter(.metric == "pr_auc") %>% 
  select(model, .config, pr_auc = mean, rank)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 01:50:30