R语言pROC包同图绘制3分类multi.roc对象多条ROC曲线
R三分类任务多模型ROC曲线同图绘制方案
三分类场景下pROC包的multiclass.roc对象仅存储成对二分类(OvO)的AUC计算结果,不直接提供可绘制的全局ROC坐标,之前三类报错的核心原因都是混淆了二分类ROC和多分类ROC的输入逻辑。以下方案采用多分类领域通用的一对其余(OvR) 策略计算每个模型的宏平均ROC,结果可解释性强,不会出现AUC数值偏差或对象提取错误。
前置准备
- 需提前安装依赖包,执行命令:
install.packages(c("pROC","ggplot2","dplyr","tidyr")) - 提前整理好两类输入数据,无需使用之前生成的
multiclass.roc对象:- 测试集真实标签向量
test_y:因子型,水平顺序与因变量Country完全一致,为c("法国","荷兰","西班牙") - 预测概率列表
model_preds:命名列表,长度为待对比的模型总数,列表名对应模型名称(如"LDA"、"梯度提升树"等);每个列表元素为三列数据框,列名与test_y的三个水平完全对应,存储测试集样本属于三个类别的预测概率。
- 测试集真实标签向量
可复现代码
步骤1:定义多分类ROC计算函数
该函数自动逐类计算二分类ROC,通过插值对齐FPR轴后生成宏平均ROC坐标,同时计算宏AUC,避免手动逐类计算的语法错误和标签错位问题。
library(pROC) library(ggplot2) library(dplyr) library(tidyr) calc_ovr_roc <- function(model_name, pred_df, true_label) { cls_levels <- levels(true_label) cls_roc <- list() cls_auc <- c() # 逐类计算单类ROC for (cls in cls_levels) { binary_label <- as.factor(ifelse(true_label == cls, 1, 0)) roc_res <- roc(response = binary_label, predictor = pred_df[[cls]], quiet = T) cls_roc[[cls]] <- data.frame( fpr = 1 - roc_res$specificities, tpr = roc_res$sensitivities ) cls_auc[cls] <- as.numeric(auc(roc_res)) } # 插值生成统一FPR网格,计算平均TPR fpr_seq <- seq(0, 1, length.out = 200) interp_tpr <- sapply(cls_roc, function(x) approx(x$fpr, x$tpr, xout = fpr_seq)$y) mean_tpr <- rowMeans(interp_tpr) # 返回标准化ROC结果 data.frame( model = model_name, fpr = fpr_seq, tpr = mean_tpr, macro_auc = round(mean(cls_auc), 3) ) }
步骤2:批量计算所有模型的ROC结果
# 遍历所有模型批量计算 roc_all <- lapply(names(model_preds), function(m) { calc_ovr_roc(m, model_preds[[m]], test_y) }) %>% bind_rows() # 生成带AUC值的图例标签 legend_labs <- roc_all %>% group_by(model) %>% summarise(auc = first(macro_auc), .groups = "drop") %>% mutate(lab = paste0(model, " (AUC = ", auc, ")")) %>% pull(lab, name = model)
步骤3:同画布绘制所有模型ROC曲线
ggplot(roc_all, aes(x = fpr, y = tpr, color = model)) + geom_line(linewidth = 0.8) + geom_abline(slope = 1, intercept = 0, linetype = "dashed", color = "gray50") + scale_x_continuous(name = "假阳性率 (1 - 特异度)", limits = c(0,1), breaks = seq(0,1,0.2)) + scale_y_continuous(name = "真阳性率 (灵敏度)", limits = c(0,1), breaks = seq(0,1,0.2)) + scale_color_discrete(name = "分类模型", labels = legend_labs) + theme_bw() + theme( panel.grid = element_blank(), legend.position = "right" )
之前三类报错的对应原因
- 逐类叠加绘图AUC值异常:手动计算时大概率出现概率列与标签水平错位、正例/负例方向颠倒的问题,上述函数强制对齐列名与因子水平,避免匹配错误
- 直接调用
roc()传入三分类值报错:pROC::roc()原生仅支持二分类响应变量,直接传入三分类标签必然触发水平数超限错误 - 提取
multiclass.roc子对象提示无有效数据:multiclass.roc存储的是3组成对二分类(法国vs荷兰、法国vs西班牙、荷兰vs西班牙)的ROC结果,并非全局平均ROC,且部分版本pROC不会完整存储ROC坐标,直接提取会返回空值
内容的提问来源于stack exchange,提问作者Alice Hobbs
相关产品推荐
相关产品推荐

