如何在workflow_map后筛选tidymodels调优参数?
解决方案
一、基于workflow_id和.config批量挑选并测试参数集
完全可以利用.config字段直接定位目标参数组合,无需手动重复定型模型,具体流程如下:
- 提取调优结果并筛选目标参数行
先取出对应workflow的调优结果,再用dplyr按多指标逻辑筛选符合要求的参数行,保留.config标识以便后续定位:
# 提取norm_nnet的调优结果 tune_results <- models_tuned %>% extract_workflow_set_result(id = "norm_nnet") # 多指标筛选示例:同时满足roc_auc≥0.9、spec≥0.8 target_params <- tune_results %>% # 先筛选roc_auc达标的.config filter(.metric == "roc_auc", .estimate >= 0.9) %>% # 再交集筛选spec达标的.config inner_join( tune_results %>% filter(.metric == "spec", .estimate >= 0.8), by = ".config" ) %>% select(.config, starts_with("penalty"), starts_with("hidden_units")) # 保留参数列与.config
- 批量生成定型Workflow并验证
用purrr批量遍历筛选出的.config,自动完成参数定型与验证集测试:
library(purrr) # 提取基础workflow模板 base_workflow <- models_tuned %>% extract_workflow(id = "norm_nnet") # 批量处理每个目标参数集 validation_results <- map_df(target_params$.config, function(config_id) { # 取出对应.config的参数组 param_set <- tune_results %>% filter(.config == config_id) %>% select_best(metric = "roc_auc") # 利用select_best提取参数结构,参数值为对应.config的结果 # 定型workflow并跑验证集,收集指标 base_workflow %>% finalize_workflow(param_set) %>% last_fit(split = split_df, metrics = mm_metrics) %>% collect_metrics() %>% mutate(.config = config_id) # 标记参数集标识 }) # 查看所有筛选参数集在验证集的表现 print(validation_results)
二、自定义多指标筛选函数
如果需要更便捷的多条件筛选,可以封装类似select_best的函数,直接返回符合多指标逻辑的参数集:
select_best_multi <- function(x, metric_filters) { # metric_filters为列表,格式如list(roc_auc = ~.x >= 0.9, spec = ~.x >= 0.8) metric_data <- collect_metrics(x) # 按指标筛选后取.config的交集 filtered_configs <- reduce( names(metric_filters), function(acc, metric) { current <- metric_data %>% filter(.metric == metric) %>% filter(eval_tidy(metric_filters[[metric]], data = .)) %>% pull(.config) intersect(acc, current) }, .init = unique(metric_data$.config) ) # 返回对应参数集 x %>% filter(.config %in% filtered_configs) %>% select_best(metric = names(metric_filters)[1]) # 可根据需求调整参数提取逻辑 } # 使用示例 multi_best_params <- select_best_multi( extract_workflow_set_result(models_tuned, id = "norm_nnet"), metric_filters = list(roc_auc = ~.x >= 0.9, spec = ~.x >= 0.8) )
三、关于tune::select_best多条件支持的说明
目前tune包的select_best确实仅支持单指标筛选,但通过上述dplyr+purrr的批量处理方式,可实现多指标筛选后的批量验证。若需要原生多指标支持,可关注tidymodels官方更新,或在其GitHub仓库提交功能需求。
内容的提问来源于stack exchange,提问作者jpm92
相关产品推荐
相关产品推荐

