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

如何基于分类结果在R的yardstick中生成pr_curve所需的多分类概率?

问题解答

1. 你的计算思路是否正确?

不正确。你当前计算的是全局范围内预测为某类的样本占比,但pr_curve需要的是每个样本针对每个类别的预测概率值(即每个样本对应所有类别的概率得分,总和通常为1)。硬分类结果(只有predicted列的类别标签)本身没有概率信息,无法直接用于绘制PR曲线——PR曲线的核心是通过调整概率阈值,观察不同阈值下的精确率和召回率变化,硬分类已经是固定阈值后的结果,没有阈值调整的空间。

2. 生成类似hpc_cv格式数据的优雅方法

场景1:基于硬分类生成“伪概率”

如果只有硬分类结果,可以将预测为某类的样本对应类别概率设为1,其余类别设为0,用dplyr和tidyr批量处理:

library(dplyr)
library(tidyr)
library(yardstick)

data <- tibble(predicted = as.factor(c("A", "A", "B", "B", "C", "C")), 
               expected = as.factor(c("A", "B", "B", "C", "A", "C")))

# 生成每个类别的伪概率列
data_with_probs <- data %>%
  mutate(prob = 1) %>%
  pivot_wider(names_from = predicted, values_from = prob, values_fill = 0) %>%
  rename_with(~paste0("pred_", .), matches("A|B|C"))

# 查看结果
data_with_probs

输出结构与hpc_cv一致,包含真实标签expected,以及每个类别的概率列(pred_A、pred_B、pred_C)。

场景2:从模型直接提取概率(推荐)

如果你的分类模型支持输出概率(比如glm、randomForest等),直接提取每个类别的概率列即可。以随机森林为例:

library(randomForest)
library(tibble)

# 假设已训练好模型
model <- randomForest(expected ~ ., data = your_training_data, probability = TRUE)
# 预测概率
pred_probs <- predict(model, newdata = your_test_data, type = "prob")
# 合并真实标签与概率列
data_with_probs <- bind_cols(your_test_data %>% select(expected), as_tibble(pred_probs))

3. 绘制多分类PR曲线

拿到符合格式的数据后,调用pr_curve时指定真实标签列和所有概率列:

# 生成所有类别作为正例的PR曲线数据
pr_curve(data_with_probs, truth = expected, pred_A, pred_B, pred_C)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 10:10:21