【问题标题】:prSummary in r caret package for imbalance datar 插入符号包中的 prSummary 用于不平衡数据
【发布时间】:2016-09-30 04:15:50
【问题描述】:

我有一个不平衡的数据,我想进行分层交叉验证并使用精确召回 auc 作为我的评估指标。

我在带有分层索引的r包插入符中使用prSummary,在计算性能时遇到错误。

以下是可以复制的样本。我发现计算 p-r auc 的样本只有 10 个,并且由于不平衡,只有一个类,因此无法计算 p-r auc。 (我发现只有十个样本来计算 p-r auc 是因为我修改了 prSummary 来强制这个函数打印出数据)

library(randomForest)
library(mlbench)
library(caret)

# Load Dataset
data(Sonar)
dataset <- Sonar
x <- dataset[,1:60]
y <- dataset[,61]
# make this data very imbalance
y[4:length(y)] <- "M"
y <- as.factor(y)
dataset$Class <- y

# create index and indexOut 
seed <- 1
set.seed(seed)
folds <- 2
idxAll <- 1:nrow(x)
cvIndex <- createFolds(factor(y), folds, returnTrain = T)
cvIndexOut <- lapply(1:length(cvIndex), function(i){
    idxAll[-cvIndex[[i]]]
})
names(cvIndexOut) <- names(cvIndex)

# set the index, indexOut and prSummaryCorrect
control <- trainControl(index = cvIndex, indexOut = cvIndexOut, 
                            method="cv", summaryFunction = prSummary, classProbs = T)
metric <- "AUC"
set.seed(seed)
mtry <- sqrt(ncol(x))
tunegrid <- expand.grid(.mtry=mtry)
rf_default <- train(Class~., data=dataset, method="rf", metric=metric, tuneGrid=tunegrid, trControl=control)

这里是错误信息:

Error in ROCR::prediction(y_pred, y_true) : 
Number of classes is not equal to 2.
ROCR currently supports only evaluation of binary classification tasks. 

【问题讨论】:

    标签: r r-caret


    【解决方案1】:

    我觉得我发现了奇怪的东西......

    即使我指定了交叉验证索引,汇总函数(无论是 prSummary 还是其他汇总函数)仍然会随机(我不确定)选择十个样本来计算性能。

    我的做法是用tryCatch定义了一个汇总函数来避免这个错误的发生。

    prSummaryCorrect <- function (data, lev = NULL, model = NULL) {
      print(data)
      print(dim(data))
      library(MLmetrics)
      library(PRROC)
      if (length(levels(data$obs)) != 2) 
        stop(levels(data$obs))
      if (length(levels(data$obs)) > 2) 
        stop(paste("Your outcome has", length(levels(data$obs)), 
                   "levels. The prSummary() function isn't appropriate."))
      if (!all(levels(data[, "pred"]) == levels(data[, "obs"]))) 
        stop("levels of observed and predicted data do not match")
    
      res <- tryCatch({
        auc <- MLmetrics::PRAUC(y_pred = data[, lev[2]], y_true = ifelse(data$obs == lev[2], 1, 0))
      }, warning = function(war) {
        print(war)
        auc <- NA
      }, error = function(e){
        print(dim(data))
        auc <- NA
      }, finally = {
        print("finally")
        auc <- NA
      })
    
      c(AUC = res,
        Precision = precision.default(data = data$pred, reference = data$obs, relevant = lev[2]), 
        Recall = recall.default(data = data$pred, reference = data$obs, relevant = lev[2]), 
        F = F_meas.default(data = data$pred, reference = data$obs, relevant = lev[2]))
    }
    

    【讨论】:

      猜你喜欢
      • 2017-02-05
      • 2020-07-15
      • 2018-04-09
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2016-12-31
      • 2018-05-01
      • 1970-01-01
      相关资源
      最近更新 更多