如何将rpart+caret二元决策树代码适配四分位数多分类变量
适配多分类(4类)的修改方案
原代码仅支持二元分类,要适配取值为1、2、3、4的多分类目标变量,需从评估函数、训练配置、交叉验证指标计算三个核心部分修改,以下是完整修改后的代码及说明:
library(rpart) library(caret) # 1. 替换二元专属的Youden指数函数,改用多分类评估函数 multiClassCustomSummary <- function(data, lev = NULL, model = NULL){ if (!all(levels(data[,"pred"]) == lev)) { stop("观测值与预测值的类别水平不匹配") } # 生成多分类混淆矩阵并提取关键指标 cm <- confusionMatrix(data[,"pred"], data[,"obs"], mode = "everything") overall_acc <- cm$overall["Accuracy"] macro_recall <- mean(cm$byClass[,"Recall"]) # 宏平均召回率,即多分类场景的平衡准确率 weighted_recall <- sum(cm$byClass[,"Recall"] * cm$byClass[,"Support"]) / sum(cm$byClass[,"Support"]) kappa <- cm$overall["Kappa"] out <- c(overall_acc, macro_recall, weighted_recall, kappa) names(out) <- c("Accuracy", "BalancedAccuracy", "WeightedRecall", "Kappa") out } # 2. 调整训练控制参数,适配多分类评估 trctrl <- trainControl(method = "repeatedcv", number = 10, repeats = 20, search = "grid", summaryFunction = multiClassCustomSummary) # 3. 训练多分类模型,选用平衡准确率作为调参依据 classifier <- train(x = training_set[, names(training_set) != "Target"], y = training_set$Target, method = 'rpart', parms = list(split = "gini"), trControl = trctrl, tuneLength = 10, metric = "BalancedAccuracy") classifier complexity_parameter <- classifier$bestTune # 4. 修改交叉验证循环,适配多分类指标计算 folds <- createFolds(dataset$Target, k = 10) cv <- lapply(folds, function(x) { training_fold <- dataset[-x, ] test_fold <- dataset[x, ] classifier <- rpart(formula = Target ~ ., data = training_fold, control = rpart.control(cp = complexity_parameter)) y_pred <- predict(classifier, newdata = test_fold[, !(names(test_fold) %in% "Target")], type = 'class') # 计算多分类混淆矩阵及指标 cm <- confusionMatrix(y_pred, test_fold$Target, mode = "everything") accuracy <- cm$overall["Accuracy"] class_recalls <- cm$byClass[,"Recall"] # 每类的召回率(对应二元分类的敏感度) balanced_accuracy <- mean(class_recalls) # 整理结果,保留总体指标与每类召回率 df <- data.frame(accuracy = accuracy, balanced_accuracy = balanced_accuracy, recall_class1 = class_recalls[1], recall_class2 = class_recalls[2], recall_class3 = class_recalls[3], recall_class4 = class_recalls[4]) return(df) }) # 计算10折交叉验证的平均指标 accuracy_avg <- Reduce("+", lapply(cv, "[[", 1))/10 balanced_accuracy_avg <- Reduce("+", lapply(cv, "[[", 2))/10 recall_class1_avg <- Reduce("+", lapply(cv, "[[", 3))/10 recall_class2_avg <- Reduce("+", lapply(cv, "[[", 4))/10 recall_class3_avg <- Reduce("+", lapply(cv, "[[", 5))/10 recall_class4_avg <- Reduce("+", lapply(cv, "[[", 6))/10 # 输出平均结果 cat("平均准确率:", accuracy_avg, "\n") cat("平均平衡准确率:", balanced_accuracy_avg, "\n") cat("类别1平均召回率:", recall_class1_avg, "\n") cat("类别2平均召回率:", recall_class2_avg, "\n") cat("类别3平均召回率:", recall_class3_avg, "\n") cat("类别4平均召回率:", recall_class4_avg, "\n")
关键修改说明:
替换二元评估逻辑:
删除原代码中仅支持二元的youdenSumary函数,改用multiClassCustomSummary,基于confusionMatrix生成多分类场景的核心指标,包括总体准确率、平衡准确率(宏平均召回率)、加权召回率和Kappa系数。调整模型训练配置:
trainControl中指定多分类评估函数,确保调参过程适配4分类场景- 选用
BalancedAccuracy作为调参指标,避免因类别不平衡导致模型偏向样本量较大的类别
重构交叉验证指标计算:
- 多分类场景下的平衡准确率定义为所有类别召回率的平均值,替代原二元逻辑中的
(敏感度+特异度)/2 - 保留每类的召回率指标,便于分析不同类别模型的表现差异
- 移除原代码中仅支持二元的混淆矩阵计算逻辑,改用通用的多分类混淆矩阵生成方式
- 多分类场景下的平衡准确率定义为所有类别召回率的平均值,替代原二元逻辑中的
内容的提问来源于stack exchange,提问作者Mark
相关产品推荐
相关产品推荐

