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

如何在tidymodels嵌套重采样中添加多评估指标?

在tidymodels嵌套重采样中添加多评估指标的实现方案

核心思路

要添加特异性、PR-AUC等指标,需要从预测结果完整性、多指标计算逻辑、结果汇总方式三个方面调整现有代码:

  • 确保预测结果包含类别预测和概率值(PR-AUC依赖概率输出)
  • 自定义函数一次性计算所有需要的评估指标
  • 调整调参和汇总逻辑,适配多指标的返回格式

修改后的完整代码

# 加载依赖包
library(mlbench)
library(tidymodels)
library(kernlab)
library(furrr)

# 生成模拟数据
sim_data <- function(n) {
  tmp <- mlbench.friedman1(n, sd = 1)
  tmp <- cbind(tmp$x, tmp$y)
  tmp <- as.data.frame(tmp)
  names(tmp)[ncol(tmp)] <- "y"
  tmp
}

set.seed(9815)
train_dat <- sim_data(50)
train_dat$y <- rep(c("yes", "no"), length.out = nrow(train_dat))
train_dat$y <- as.factor(train_dat$y)

# 设置嵌套交叉验证
results <- nested_cv(train_dat, 
                     outside = vfold_cv(v= 3, repeats = 3), 
                     inside = bootstraps(times = 5))

# 1. 自定义多指标计算函数:同时计算灵敏度、特异性、PR-AUC
svm_metrics <- function(object, cost = 1, rbf_sigma = 0.2) {
  # 训练SVM模型
  mod <- 
    svm_rbf(mode = "classification", cost = cost, rbf_sigma = rbf_sigma) %>% 
    set_engine("kernlab") %>% 
    fit(y ~ ., data = analysis(object))
  
  # 获取完整预测结果:类别预测 + 概率预测
  holdout_pred <- 
    predict(mod, assessment(object), type = "class") %>%  # 类别预测结果
    bind_cols(predict(mod, assessment(object), type = "prob")) %>%  # 概率预测结果
    bind_cols(assessment(object) %>% dplyr::select(y))  # 真实标签
  
  # 计算所有目标指标
  tibble(
    sens = sens(holdout_pred, truth = y, estimate = .pred_class)$.estimate,
    spec = spec(holdout_pred, truth = y, estimate = .pred_class)$.estimate,
    pr_auc = pr_auc(holdout_pred, truth = y, .pred_yes)$.estimate  # 指定正类概率列
  )
}

# 2. 参数包装函数:适配多指标返回格式
svm_metrics_wrapper <- function(cost, rbf_sigma, object) {
  svm_metrics(object, cost, rbf_sigma)
}

# 3. 调参循环:处理多指标结果
tune_over_svm <- function(object){
  tibble(cost = grid_random(cost(), size = 3),
         rbf_sigma = grid_random(rbf_sigma(), size = 3)) %>% 
    mutate(metrics = map2(cost, rbf_sigma, svm_metrics_wrapper, object = object)) %>% 
    unnest(metrics)  # 展开多指标列,方便后续汇总
}

# 4. 汇总调参结果:按超参数分组计算各指标均值
summarize_tune_results <- function(object) {
  map_df(object$splits, tune_over_svm) %>%
    group_by(cost, rbf_sigma) %>%
    summarize(
      mean_sens = mean(sens, na.rm = TRUE),
      mean_spec = mean(spec, na.rm = TRUE),
      mean_pr_auc = mean(pr_auc, na.rm = TRUE),
      n = n(),
      .groups = "drop"
    )
}

# 并行计算调参结果
plan(multisession)
tuning_results <- future_map(results$inner_resamples, summarize_tune_results) 

关键修改说明

  • 预测结果扩展:新增type = "prob"获取概率值,PR-AUC指标需要基于正类的预测概率计算,需根据你的标签(如yes)指定对应的概率列(.pred_yes)。
  • 多指标函数:将原单一返回灵敏度的函数,改为返回包含所有指标的tibble,后续可轻松添加更多指标(如roc_auc、accuracy等)。
  • 调参逻辑适配:用map2()替代map2_dbl(),因为现在返回的是多值结果,再通过unnest()展开成列,方便分组汇总。
  • 汇总逻辑更新:在summarize中对每个指标分别计算均值,保留超参数分组逻辑。

内容的提问来源于stack exchange,提问作者Tengku Hanis

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 09:06:30