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

