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

