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

在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)
}

额外验证点

  1. 若二分类标签的正类不是默认的第二个水平,需在metric_set中指定event_level:
metrics = metric_set(pg, event_level = "first") # 或"second",根据实际标签水平调整
  1. 单独验证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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 13:13:24