You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将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")

关键修改说明:

  1. 替换二元评估逻辑:
    删除原代码中仅支持二元的youdenSumary函数,改用multiClassCustomSummary,基于confusionMatrix生成多分类场景的核心指标,包括总体准确率、平衡准确率(宏平均召回率)、加权召回率和Kappa系数。

  2. 调整模型训练配置:

    • trainControl中指定多分类评估函数,确保调参过程适配4分类场景
    • 选用BalancedAccuracy作为调参指标,避免因类别不平衡导致模型偏向样本量较大的类别
  3. 重构交叉验证指标计算:

    • 多分类场景下的平衡准确率定义为所有类别召回率的平均值,替代原二元逻辑中的(敏感度+特异度)/2
    • 保留每类的召回率指标,便于分析不同类别模型的表现差异
    • 移除原代码中仅支持二元的混淆矩阵计算逻辑,改用通用的多分类混淆矩阵生成方式

内容的提问来源于stack exchange,提问作者Mark

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.23 23:27:49