在tidymodels中使用自定义概率指标pg时出现错误
自定义部分基尼系数(pg)在tidymodels调参中的报错解决方法
问题场景
使用tidymodels对XGBoost模型调参时,将自定义部分基尼系数(pg)作为评价指标,代码在使用roc_auc时正常运行,但切换为pg后出现所有模型失败的警告,报错信息显示no applicable method for 'pg' applied to an object of class "c('grouped_df', 'tbl_df', 'tbl', 'data.frame')"。
调参代码:
xgb_folds <- train %>% vfold_cv(v=5) xgb_model <- parsnip::boost_tree( mode = "classification", trees = tune(), tree_depth = tune(), learn_rate = tune(), loss_reduction = tune() ) %>% set_engine("xgboost") xgb_wf <- workflow() %>% add_recipe(TREE_recipe) %>% add_model(xgb_model) xgboost_tuned <- tune::tune_grid( object = xgb_wf, resamples = xgb_folds, grid = hyperparameters_XGB_tidy, metrics = metric_set(pg), control = tune::control_grid(verbose = TRUE) )
错误信息:
unique notes: ──────────────────────────────────────────────────────────────────────────────────────────────────── Error in `metric_set()`: ! Failed to compute `pg()`. Caused by error in `UseMethod()`: ! no applicable method for 'pg' applied to an object of class "c('grouped_df', 'tbl_df', 'tbl', 'data.frame')"
自定义pg指标实现代码:
# partialGini for tidymodels library(tidymodels, rlang) pg_impl <- function(truth, estimate, case_weights = NULL) { sorted_indices <- order(estimate, decreasing = TRUE) sorted_probs <- estimate[sorted_indices] sorted_actuals <- truth[sorted_indices] # Select subset with PD < 0.4 subset_indices <- which(sorted_probs < 0.4) subset_probs <- sorted_probs[subset_indices] subset_actuals <- sorted_actuals[subset_indices] # Check if there are both positive and negative cases in the subset if (length(unique(subset_actuals)) > 1) { # Calculate ROC curve for the subset roc_subset <- pROC::roc(subset_actuals, subset_probs, direction = "<", quiet = TRUE) # Calculate AUC for the subset partial_auc <- pROC::auc(roc_subset) # Calculate partial Gini coefficient (2 * partial_auc - 1) } else return(NA) } pg_vec <- function(truth, estimate, estimator = NULL, na_rm = TRUE, case_weights = NULL, ...) { abort_if_class_pred(truth) estimator <- finalize_estimator(truth, estimator) check_prob_metric(truth, estimate, case_weights, estimator) if (na_rm) { result <- yardstick_remove_missing(truth, estimate, case_weights) truth <- result$truth estimate <- result$estimate case_weights <- result$case_weights } else if (yardstick_any_missing(truth, estimate, case_weights)) { return(NA_real_) } pg_impl(truth, estimate, case_weights = case_weights) } pg <- function(data, ...) { UseMethod("pg") } pg <- new_prob_metric(pg, direction = "maximize") pg.data.frame <- function(data, truth, ..., na_rm = TRUE) { prob_metric_summarizer( name = "pg", fn = pg_vec, data = data, truth = !! enquo(truth), ..., na_rm = na_rm) }
问题原因
交叉验证生成的resamples数据是分组数据框(grouped_df),但自定义的pg指标仅实现了data.frame类型的处理方法,R的方法分派机制找不到对应grouped_df的处理逻辑,因此报错。
解决方案
方案1:添加grouped_df专属处理方法
在自定义pg指标的代码末尾,添加以下代码:
pg.grouped_df <- function(data, truth, ..., na_rm = TRUE) { # 先取消分组,再调用已实现的data.frame方法 data %>% dplyr::ungroup() %>% pg.data.frame(truth = {{truth}}, ..., na_rm = na_rm) }
方案2:修改主函数自动适配分组
替换原有的pg主函数定义:
# 替换原来的 pg <- function(data, ...) { UseMethod("pg") } pg <- function(data, ...) { # 若为分组数据框,先取消分组 if (inherits(data, "grouped_df")) { data <- dplyr::ungroup(data) } UseMethod("pg", data) }
额外验证点
- 若二分类标签的正类不是默认的第二个水平,需在
metric_set中指定event_level:
metrics = metric_set(pg, event_level = "first") # 或"second",根据实际标签水平调整
- 单独验证
pg函数对分组数据框的处理能力,确保逻辑正常:
# 构造测试分组数据 test_data <- tibble( y = factor(c(1,0,1,0,1,0)), .pred_1 = c(0.6, 0.3, 0.5, 0.2, 0.7, 0.1) ) %>% group_by(y) # 测试pg函数 pg(test_data, truth = y, .pred_1)
内容的提问来源于stack exchange,提问作者Simon De Lange
相关产品推荐
相关产品推荐

