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

在tidymodels中为ROC-AUC与Brier分数设置案例权重

解决tidymodels中加权性能指标计算的问题

你已经通过add_case_weights()在工作流中加入了样本权重,但默认的roc_auc和brier_class指标并不会自动使用这些权重,需要自定义支持权重的指标函数来实现需求。

步骤1:自定义加权指标函数

创建支持样本权重的ROC-AUC和Brier评分函数:

# 加载所需包
library(pROC)
library(yardstick)

# 自定义加权ROC-AUC指标
weighted_roc_auc <- function(data, truth, estimate, weights, ...) {
  # 提取真实标签、预测概率和权重
  y_true <- data[[truth]]
  y_prob <- data[[estimate]]
  w <- data[[weights]]
  
  # 使用pROC计算加权ROC-AUC
  roc_obj <- roc(y_true, y_prob, weights = w, quiet = TRUE)
  auc(roc_obj)
}

# 注册为yardstick指标,设置二分类模式
weighted_roc_auc <- metric_summarizer(
  metric_nm = "weighted_roc_auc",
  metric_fn = weighted_roc_auc,
  direction = "maximize",
  truth = "truth",
  estimate = "estimate",
  weights = "weights",
  mode = "classification"
)

# 自定义加权Brier评分(二分类)
weighted_brier_class <- function(data, truth, estimate, weights, ...) {
  y_true <- as.numeric(data[[truth]]) - 1  # 转为0/1数值
  y_prob <- data[[estimate]]
  w <- data[[weights]]
  
  # 计算加权平均的平方误差
  mean(w * (y_true - y_prob)^2)
}

# 注册为yardstick指标
weighted_brier_class <- metric_summarizer(
  metric_nm = "weighted_brier_class",
  metric_fn = weighted_brier_class,
  direction = "minimize",
  truth = "truth",
  estimate = "estimate",
  weights = "weights",
  mode = "classification"
)

步骤2:使用自定义指标进行交叉验证

将自定义指标用metric_set()包装,传入fit_resamples()的metrics参数:

# 创建加权指标集合
weighted_metrics <- metric_set(weighted_roc_auc, weighted_brier_class)

# 重新运行交叉验证
current_fit <- fit_resamples(
  current_wf,
  resamples = data_folds,
  metrics = weighted_metrics,
  control = cntrl
)

# 查看结果
collect_metrics(current_fit)

关键说明

  • 工作流中通过add_case_weights(imp_weights)添加的权重,会自动传递给自定义指标函数的weights参数。
  • 对于group_vfold_cv的分组交叉验证,每个折叠内的权重都会被正确应用,无需额外处理。
  • 如果你的Status30D是因子类型,自定义函数中as.numeric(data[[truth]]) - 1会将其转为0/1数值,确保Brier评分计算正确。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 21:54:59