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

如何基于workflow_map结果高效拟合并预测每个模型?

批量拟合workflow_map调参后的所有超参数模型并预测

针对批量处理所有调参组合的需求,我们可以借助purrr迭代工具结合tidymodels原生函数,高效完成所有超参数模型的拟合与预测,无需手动逐个操作。

步骤1:提取调参网格与基础工作流

首先从wf_tuning中拆分出每个模型的所有超参数组合,以及对应的基础工作流模板:

library(purrr)
library(dplyr)
library(tidyr)

# 提取每个模型的调参结果和基础工作流
model_metadata <- wf_tuning %>%
  mutate(
    # 提取所有超参数组合(自动适配不同模型的参数名)
    param_grids = map(result, ~ .x$.metrics %>% 
                        select(-.metric, -.estimator, -.mean, -.n, -.std_err) %>% 
                        distinct()),
    # 提取对应模型的基础工作流
    base_workflow = map(wflow_id, ~ extract_workflow(wf_tuning, .x))
  )

步骤2:批量更新、拟合与预测

定义一个处理单模型所有超参数组合的函数,然后批量应用到所有模型:

# 处理单个模型的所有超参数组合
process_single_model <- function(base_wf, param_grid) {
  param_grid %>%
    mutate(
      # 用当前超参数组合更新工作流
      updated_workflow = pmap(., function(...) {
        param_list <- list(...)
        base_wf %>% update_model(parameters(param_list))
      }),
      # 在训练集拟合模型
      fitted_workflow = map(updated_workflow, ~ fit(.x, data = TrainData)),
      # 对测试集生成分类预测
      test_predictions = map(fitted_workflow, ~ predict(.x, new_data = TestData, type = "class")),
      # 合并真实标签与预测结果
      results = map2(test_predictions, list(TestData$Y), ~ tibble(truth = .y, prediction = .x$.pred_class))
    )
}

# 批量处理所有模型
all_model_results <- model_metadata %>%
  mutate(
    model_output = map2(base_workflow, param_grids, process_single_model)
  )

# 整合所有结果为一个整洁的数据框
final_combined_results <- all_model_results %>%
  select(wflow_id, model_output) %>%
  unnest(model_output) %>%
  unnest(results)

步骤3:识别未预测案例

如果存在模型拟合失败导致的无预测结果,可以通过筛选缺失值定位:

# 查找所有未生成预测的案例
unpredicted_cases <- final_combined_results %>%
  filter(is.na(prediction)) %>%
  # 关联测试集原始数据,查看具体案例信息
  bind_cols(TestData[rep(1:nrow(TestData), nrow(param_grid)), ]) # 根据实际数据结构调整关联逻辑

优化建议

  • 并行加速:如果模型数量多、参数组合多,用furrr包开启并行计算:
    library(furrr)
    plan(multisession) # 根据CPU核心数设置并行进程
    all_model_results <- model_metadata %>%
      mutate(
        model_output = future_map2(base_workflow, param_grids, process_single_model)
      )
    
  • 参数适配:上述代码通过排除指标列的方式自动适配不同模型的参数,无需手动指定参数名,通用性更强。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 11:53:12