如何在R语言Caret包多分类任务中使用F1 Score作为评估指标
问题原因
报错触发原因
- 你当前使用的
posPredValue()、sensitivity()均为二分类场景专属函数,传入多分类的预测标签、真实标签时,就会触发「输入必须只有两个水平」的报错。 - 你在
trainControl配置中summaryFunction参数仍设为默认的defaultSummary,未替换为自定义的F1计算函数,即使函数逻辑无误也不会生效。
自定义函数中data仅10行的原因
caret的summaryFunction接收的data参数为当前交叉验证折对应的验证集预测结果,你使用的data.small本身样本量较小,加上配置了2折交叉验证,每折拆分出的验证集样本量自然很小,就会出现仅10行的情况。
解决方案
重写支持多分类的F1计算summaryFunction,替换原有二分类逻辑即可,以下是宏平均F1的实现示例:
# 多分类宏F1自定义评估函数 multiClassF1 <- function(data, lev = NULL, model = NULL) { cm <- caret::confusionMatrix(data$pred, data$obs) # 提取所有类别的精确率、召回率 precision <- cm[["byClass"]][, "Precision"] recall <- cm[["byClass"]][, "Recall"] # 单类F1计算后取平均得到宏F1 f1_val <- mean((2 * precision * recall)/(precision + recall), na.rm = TRUE) names(f1_val) <- "F1" return(f1_val) }
修改trainControl配置,替换summaryFunction为上述自定义函数:
train.control <- trainControl(method = "repeatedcv", number = 2, summaryFunction = multiClassF1, classProbs = TRUE, search = "grid")
后续训练代码无需修改,即可正常以多分类F1为指标训练模型。如果需要使用微平均F1,只需将上述计算逻辑替换为全局TP/FP/FN的统计即可。
内容的提问来源于stack exchange,提问作者YannickAaron
相关产品推荐
相关产品推荐

