Caret计算F1分数的疑问:为何结果与理论公式不符?
问题原因与解决方案
核心问题是你在二分类任务中错误使用了多分类评估函数multiClassSummary,导致输出的F1、Precision、Recall并非你预期的单个类别的指标:
multiClassSummary是为多分类设计的,它返回的F1是所有类别的macro平均F1(即每个类别的F1取算术平均值),而非单个类别的F1。- 你看到的
Precision和Recall同样是所有类别的算术平均值,不是单个类别的指标,因此代入二分类F1公式计算自然会和输出的F1不符。
你的结果中F1=0.625,说明另一个类别的F1值很高(比如接近1),两者平均后得到这个结果——而你手动计算的只是少数类的F1(0.105),和macro平均的F1完全不是同一个指标。
解决方法
方法1:改用二分类专用评估函数
将trainControl中的summaryFunction替换为twoClassSummary,同时确保你的响应变量y是二分类因子(如果不是,先转换):
# 确保y是二分类因子,根据实际类别名调整levels y <- factor(y, levels = c("neg", "pos")) control <- trainControl(method="repeatedcv", number=k_folds_cross_validation, repeats=times_cross_validation, search = "random", summaryFunction = twoClassSummary, returnResamp = "all", savePredictions = "all", classProbs = TRUE) set.seed(7) tuneGrid = expand.grid(alpha = 0, lambda = 10^seq(-4, 0, length.out = 50)) lasso_ML <- train(t(x), y, method = "glmnet", trControl = control, family="binomial", tuneGrid = tuneGrid, metric="ROC")
默认twoClassSummary不直接返回F1,如果需要正类的F1,可自定义评估函数:
# 自定义二分类F1评估函数 twoClassF1 <- function(data, lev = NULL, model = NULL) { if (!all(levels(data$pred) == lev)) { stop("观测值与预测值的类别不匹配") } # 计算混淆矩阵,指定正类为lev[2] cm <- confusionMatrix(data$pred, data$obs, positive = lev[2]) precision <- cm$byClass["Precision"] recall <- cm$byClass["Recall"] f1 <- 2 * (precision * recall) / (precision + recall) # 返回结果 out <- c(F1 = unname(f1), Precision = unname(precision), Recall = unname(recall)) out } # 在trainControl中使用自定义函数 control <- trainControl(method="repeatedcv", number=k_folds_cross_validation, repeats=times_cross_validation, search = "random", summaryFunction = twoClassF1, returnResamp = "all", savePredictions = "all", classProbs = TRUE) # 训练模型时指定metric="F1" lasso_ML <- train(t(x), y, method = "glmnet", trControl = control, family="binomial", tuneGrid = tuneGrid, metric="F1")
方法2:手动计算单个类别的F1
如果你坚持使用multiClassSummary,可以从模型的预测结果中提取单个类别的指标:
# 提取所有交叉验证的预测结果 preds <- lasso_ML$pred # 按lambda分组计算每个组的混淆矩阵和F1 library(dplyr) grouped_metrics <- preds %>% group_by(lambda) %>% summarise( # 替换为你的正类名称 TP = sum(pred == "pos" & obs == "pos"), FP = sum(pred == "pos" & obs != "pos"), FN = sum(pred != "pos" & obs == "pos"), TN = sum(pred != "pos" & obs != "pos"), # 计算Precision, Recall, F1 Precision = TP / (TP + FP), Recall = TP / (TP + FN), F1 = 2 * (Precision * Recall) / (Precision + Recall) ) # 查看结果 print(grouped_metrics)
内容的提问来源于stack exchange,提问作者Lluc
相关产品推荐
相关产品推荐

