在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
相关产品推荐
相关产品推荐

