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

如何用tidyverse/tidymodels或base R按正例百分位数调阈值计算二分类指标?

基于tidyverse/tidymodels的百分位数阈值调整方法

针对你需要通过设定正例百分位数来调整分类阈值、计算二分类指标的需求,以下是几种稳健的实现方案:

1. 核心思路

要让约N%的样本被预测为正例,本质是找到一个概率阈值k,使得大于等于k的正例概率样本占比约为N%。可以直接利用分位数函数或百分位排序来定位这个阈值,再结合tidymodels的工具计算分类指标。

2. tidyverse + tidymodels实现

步骤1:确定百分位数阈值

使用quantile()函数提取正例概率的对应分位数(注意:需取1-目标占比的分位数,比如要20%正例,取0.8分位数,因为前80%的概率低于该阈值,剩余20%高于阈值会被判定为正例):

library(tidyverse)
library(tidymodels)

# 假设你的预测结果数据框为pred_df
pred_df <- rf_fit %>% 
  predict(test, type = "prob") %>% 
  bind_cols(test %>% select(Day90))

# 设定目标正例占比:20%
target_pct <- 0.2
# 计算阈值
threshold <- quantile(pred_df$.pred_1, 1 - target_pct)

# 生成预测类别
pred_df <- pred_df %>% 
  mutate(.pred_class = factor(ifelse(.pred_1 >= threshold, "1", "0")))

步骤2:计算分类指标

用tidymodels的metrics()函数一键生成常用指标:

# 计算多类指标(含精度、召回率、ROC-AUC等)
metrics(pred_df, truth = Day90, estimate = .pred_class, prob = .pred_1)

# 单独提取特定指标,比如ROC-AUC
roc_auc(pred_df, truth = Day90, .pred_1)

3. 用dplyr::percent_rank()处理重复概率

如果遇到大量重复的正例概率值,percent_rank()可以更灵活地匹配目标百分位:

pred_df <- pred_df %>% 
  mutate(pct_rank = percent_rank(.pred_1))

# 找到刚好让约20%样本为正例的阈值:取百分位≥0.8的最小概率值
threshold <- pred_df %>% 
  filter(pct_rank >= (1 - target_pct)) %>% 
  pull(.pred_1) %>% 
  min()

# 生成预测类别并计算指标
pred_df <- pred_df %>% 
  mutate(.pred_class = factor(ifelse(.pred_1 >= threshold, "1", "0")))

metrics(pred_df, truth = Day90, estimate = .pred_class)

4. Base R实现方案

如果不用tidyverse,Base R也能快速完成:

# 提取正例概率
pred_probs <- pred_df$.pred_1
# 计算阈值
threshold <- quantile(pred_probs, 1 - target_pct)
# 生成预测类别
pred_class <- factor(ifelse(pred_probs >= threshold, "1", "0"))
# 计算混淆矩阵与核心指标
conf_mat <- table(pred_class, pred_df$Day90)
accuracy <- sum(diag(conf_mat)) / sum(conf_mat)
recall <- conf_mat[2,2] / sum(conf_mat[,2])
precision <- conf_mat[2,2] / sum(conf_mat[2,])

注意事项

  • 当存在重复概率时,quantile()的type参数(可选1-9)会影响阈值结果,type=6更贴合直观的百分位定义;
  • 如果需要精确的N%正例(而非近似值),可直接对概率降序排序后取第nrow(pred_df)*target_pct位的概率:
sorted_probs <- sort(pred_df$.pred_1, decreasing = TRUE)
threshold <- sorted_probs[floor(nrow(pred_df)*target_pct)]

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.03 17:56:17