【问题标题】:ROC curve for Training set and Test set for each fold of cross validation in CaretCaret 中每个交叉验证折叠的训练集和测试集的 ROC 曲线
【发布时间】:2018-04-03 22:04:25
【问题描述】:

是否可以为 Caret 中的 5 折交叉验证中的每一折分别设置训练集和测试集的 ROC 曲线?

library(caret)
train_control <- trainControl(method="cv", number=5,savePredictions =  TRUE,classProbs = TRUE)
output <- train(Species~., data=iris, trControl=train_control, method="rf")

我可以执行以下操作,但我不知道它是否会为 Fold1 的训练集或测试集返回 ROC:

library(pROC) 
selectedIndices <- rfmodel$pred$Resample == "Fold1"
plot.roc(rfmodel$pred$obs[selectedIndices],rfmodel$pred$setosa[selectedIndices])

【问题讨论】:

    标签: r machine-learning cross-validation r-caret roc


    【解决方案1】:

    确实,documentationrfmodel$pred 的内容一点也不清楚——我敢打赌,所包含的预测是针对用作测试集的折叠,但我不能指出其中的任何证据文档;尽管如此,不管怎样,您在尝试获得 ROC 的过程中仍然缺少一些要点。

    首先,让我们将rfmodel$pred 隔离在一个单独的数据框中以便于处理:

    dd <- rfmodel$pred
    
    nrow(dd)
    # 450
    

    为什么是 450 行?这是因为您已经尝试了 3 个不同的参数集(在您的情况下,mtry 只使用了 3 个不同的值):

    rfmodel$results
    # output:
      mtry Accuracy Kappa AccuracySD    KappaSD
    1    2     0.96  0.94 0.04346135 0.06519202
    2    3     0.96  0.94 0.04346135 0.06519202
    3    4     0.96  0.94 0.04346135 0.06519202
    

    150 行 X 3 设置 = 450。

    让我们仔细看看rfmodel$pred的内容:

    head(dd)
    
    # result:
        pred    obs setosa versicolor virginica rowIndex mtry Resample
    1 setosa setosa  1.000      0.000         0        2    2    Fold1
    2 setosa setosa  1.000      0.000         0        3    2    Fold1
    3 setosa setosa  1.000      0.000         0        6    2    Fold1
    4 setosa setosa  0.998      0.002         0       24    2    Fold1
    5 setosa setosa  1.000      0.000         0       33    2    Fold1
    6 setosa setosa  1.000      0.000         0       38    2    Fold1
    
    • obs 列包含真实值
    • setosaversicolorvirginica 三列分别包含为每个类计算的概率,它们每行的总和为 1
    • pred 列包含最终预测,即上述三列中概率最大的类

    如果这就是整个故事,那么您绘制 ROC 的方式就可以了,即:

    selectedIndices <- rfmodel$pred$Resample == "Fold1"
    plot.roc(rfmodel$pred$obs[selectedIndices],rfmodel$pred$setosa[selectedIndices])
    

    但这并不是故事的全部(仅仅存在 450 行而不是 150 行应该已经给出了提示):请注意名为 mtry 的列的存在;事实上,rfmodel$pred 包含了所有次交叉验证运行的结果(即所有参数设置):

    tail(dd)
    # result:
             pred       obs setosa versicolor virginica rowIndex mtry Resample
    445 virginica virginica      0      0.004     0.996      112    4    Fold5
    446 virginica virginica      0      0.000     1.000      113    4    Fold5
    447 virginica virginica      0      0.020     0.980      115    4    Fold5
    448 virginica virginica      0      0.000     1.000      118    4    Fold5
    449 virginica virginica      0      0.394     0.606      135    4    Fold5
    450 virginica virginica      0      0.000     1.000      140    4    Fold5
    

    这就是你的selectedIndices计算不正确的根本原因;它还应该包括mtry 的特定选择,否则 ROC 没有任何意义,因为它“聚合”了多个模型:

    selectedIndices <- rfmodel$pred$Resample == "Fold1" & rfmodel$pred$mtry == 2
    

    --

    正如我一开始所说,我敢打赌rfmodel$pred 中的预测是针对该文件夹作为测试集的;事实上,如果我们手动计算准确度,它们与上面显示的rfmodel$results 中报告的准确度一致(所有 3 种设置均为 0.96),我们知道这是用作 test 的文件夹(可以说,各自的训练精度为 1.0):

    for (i in 2:4) {  # mtry values in {2, 3, 4}
    
    acc = (length(which(dd$pred == dd$obs & dd$mtry==i & dd$Resample=='Fold1'))/30 +
        length(which(dd$pred == dd$obs & dd$mtry==i & dd$Resample=='Fold2'))/30 +
        length(which(dd$pred == dd$obs & dd$mtry==i & dd$Resample=='Fold3'))/30 +
        length(which(dd$pred == dd$obs & dd$mtry==i & dd$Resample=='Fold4'))/30 +
        length(which(dd$pred == dd$obs & dd$mtry==i & dd$Resample=='Fold5'))/30
    )/5
    
    print(acc) 
    }
    
    # result:
    [1] 0.96
    [1] 0.96
    [1] 0.96
    

    【讨论】:

      猜你喜欢
      • 2018-04-03
      • 2019-01-25
      • 2016-05-12
      • 2020-10-30
      • 2021-09-14
      • 2020-01-02
      • 2016-09-09
      • 2018-05-03
      • 2012-09-11
      相关资源
      最近更新 更多