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

如何在yardstick中创建平衡对数损失函数用于tidymodels工作流?

在yardstick中创建平衡对数损失函数并用于tidymodels工作流

平衡对数损失(Balanced Log Loss)是针对类别不平衡场景优化的损失指标,核心是对正负样本的对数损失分别计算后按样本量加权平均。下面是具体实现步骤:

1. 二分类场景的平衡对数损失公式

平衡对数损失 = - (1/N₊) * Σ(yᵢ * log(pᵢ)) - (1/N₋) * Σ((1 - yᵢ) * log(1 - pᵢ))
其中:

  • N₊:正样本数量
  • N₋:负样本数量
  • yᵢ:第i个样本的真实标签(1为正类,0为负类)
  • pᵢ:第i个样本被预测为正类的概率

2. 自定义二分类平衡对数损失指标

使用yardstick的new_prob_metric()函数创建兼容tidymodels的指标,同时编写计算逻辑:

library(yardstick)
library(dplyr)
library(purrr)

# 定义二分类平衡对数损失的向量计算函数
balanced_log_loss_vec <- function(truth, estimate, event_level = "first", ...) {
  # 统一真实标签的格式,处理事件类别位置
  truth <- if (event_level == "second") factor(truth, levels = rev(levels(truth))) else factor(truth)
  truth_bin <- as.integer(truth) - 1
  
  # 提取正类的预测概率(兼容矩阵或向量输入)
  p_pos <- if (is.matrix(estimate)) estimate[, levels(truth)[2]] else estimate
  
  # 避免log(0)或log(1)的极端情况,添加微小偏移
  p_pos <- pmax(pmin(p_pos, 1 - 1e-15), 1e-15)
  
  # 拆分正负样本组
  pos_idx <- truth_bin == 1
  neg_idx <- !pos_idx
  
  n_pos <- sum(pos_idx)
  n_neg <- sum(neg_idx)
  
  # 分别计算正负样本的对数损失
  loss_pos <- if (n_pos > 0) -mean(log(p_pos[pos_idx])) else 0
  loss_neg <- if (n_neg > 0) -mean(log(1 - p_pos[neg_idx])) else 0
  
  # 返回平衡后的损失值
  (loss_pos + loss_neg) / 2
}

# 创建yardstick兼容的指标对象
balanced_log_loss <- new_prob_metric(
  balanced_log_loss_vec,
  direction = "minimize",  # 损失越小模型性能越好
  type = "binary"
)

3. 测试自定义指标

用yardstick自带的two_class_example数据集验证:

data(two_class_example)

# 计算平衡对数损失
balanced_log_loss(two_class_example, truth, .pred_Class1)
# 输出示例:[1] 0.3245...

4. 扩展到多分类场景

如果需要支持多分类,调整计算逻辑,对每个类别计算对数损失后按类别样本量加权:

# 多分类平衡对数损失的向量计算函数
balanced_log_loss_multi_vec <- function(truth, estimate, ...) {
  truth <- factor(truth)
  classes <- levels(truth)
  
  # 将真实标签转换为one-hot编码
  truth_onehot <- model.matrix(~ truth - 1, data = data.frame(truth))
  
  # 避免极端概率值
  estimate <- pmax(pmin(estimate, 1 - 1e-15), 1e-15)
  
  # 计算每个类别的对数损失
  class_losses <- map_dbl(seq_along(classes), function(i) {
    class_idx <- truth == classes[i]
    n_class <- sum(class_idx)
    if (n_class == 0) return(0)
    -mean(log(estimate[class_idx, classes[i]]))
  })
  
  # 按每个类别的样本占比加权平均
  class_weights <- table(truth) / nrow(estimate)
  sum(class_losses * class_weights)
}

# 创建多分类版本的指标对象
balanced_log_loss_multi <- new_prob_metric(
  balanced_log_loss_multi_vec,
  direction = "minimize",
  type = "multiclass"
)

5. 在tidymodels工作流中使用

在交叉验证或最终拟合时,将自定义指标加入metrics参数即可:

library(tidymodels)

# 示例:构建逻辑回归工作流
lr_spec <- logistic_reg() %>% set_engine("glm") %>% set_mode("classification")
wf <- workflow() %>% add_model(lr_spec) %>% add_formula(Class ~ .)

# 交叉验证时使用自定义指标
set.seed(123)
cv_res <- fit_resamples(
  wf,
  resamples = vfold_cv(two_class_example, v = 5),
  metrics = metric_set(balanced_log_loss, roc_auc, accuracy)
)

# 查看交叉验证结果
collect_metrics(cv_res)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 08:03:12