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

如何在workflow_map后筛选tidymodels调优参数?

解决方案

一、基于workflow_id和.config批量挑选并测试参数集

完全可以利用.config字段直接定位目标参数组合,无需手动重复定型模型,具体流程如下:

  1. 提取调优结果并筛选目标参数行
    先取出对应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
  1. 批量生成定型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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 02:15:31