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

关于R语言gramEvol包符号回归实现正确性的技术问询

问题描述

我使用R语言的gramEvol包进行符号回归,在IRIS数据集上测试时运行正常,但在包含大量变量且存在类别不平衡的其他数据集上,发现实现似乎存在问题。怀疑SymRegFitFunc函数或ruleDef存在实现错误,请帮忙验证以下IRIS数据集示例代码的实现是否正确:

# Load the required libraries
#### Setup ###
source("_Setup.R")
library(magrittr)
library(gramEvol)

kFolds = 10L
resultados <- list()
resultados$Desempenhos <- tibble()

# Load the iris dataset
data(iris)

dadosModelo <- iris %>%
  mutate(Species = as.factor(Species) %>% as.numeric()) %>%
  mutate(Species = ifelse(test = Species == 1, yes = 0,1)) %>% 
  rename_with(~ str_c("x", .x), -Species) %>% 
  select(y = Species, everything()) %>% 
  as_tibble()


nVars = ncol(dadosModelo) - 1
nomeBase <- "Iris"

#### Split dataset  ####
resultados$DivisaoDados[[nomeBase]] <- SplitDataset(
  dados = dadosModelo,
  kFolds = kFolds, 
  porcaoTreinoKS = 0.9, 
  arquivoCache = NULL,  
  buscarCache = FALSE
)

linhasTreinoTeste = resultados$DivisaoDados[[nomeBase]]


for (fold in 1:kFolds) {
  
  Msg(
    bold("Base de Dados: "), 
    bold(yellow(nomeBase)),
    " ___ ",
    "/",
    "\n \n",
    bold(blue("Fold: ")),
    bold(blue(fold)),
    "/",
    bold(kFolds)
  )
  
  # Extracts training and testing data from folds
  dadosTreinoCV <- dadosModelo[linhasTreinoTeste$LinhasTreinoTesteFolds$LinhasTreino[[fold]] %>% unlist(), ]
  dadosTesteCV <- dadosModelo[linhasTreinoTeste$LinhasTreinoTesteFolds$LinhasTeste[[fold]] %>% unlist(), ]
  
  #### Symbolic Regression ####
  # Training data
  x_train <- dadosTreinoCV[, 2:nVars]
  y_train <- dadosTreinoCV$y
  
  # Testing data
  x_test <- dadosTesteCV[, 2:nVars]
  y_test <- dadosTesteCV$y
  
  # Generate the list of variables dynamically
  var_names <- colnames(x_train)
  
  # Create the grammar rules for symbolic regression
  var_rule <- do.call(grule, lapply(var_names, as.name))
  
  ruleDef <- list(
    expr = grule(op(expr, expr), func(expr), var),
    func = grule(sin, cos, sqrt, log, exp), # grule(sin, cos, log, exp),
    op = grule('+', '-', '*', '/'),
    var = var_rule
  )
  
  # Create grammar from defined rules
  grammarDef <- CreateGrammar(ruleDef)
  
  # Tuning function 
  SymRegFitFunc <- function(expr) {
    env <- as.list(x_train)
    result <- try(eval(expr, envir = env), silent = TRUE)
    
    if (inherits(result, "try-error") || is.language(result) || length(result) != length(y_train)) {
      return(Inf)  # Usar Inf para indicar um ajuste ruim
    } else {
      
      # Converter previsões contínuas para valores discretos (categorias)
      predicted_classes <- ifelse(result < 0.5, 0, 1)
      # predicted_classes <- ifelse(result < 2, 0, 1)
      
      # Calcular a matriz de confusão
      confusion_matrix <- table(predicted_classes, y_train)
      
      
      # Verdadeiros positivos
      VP <- ifelse(
        test = ("1" %in% colnames(confusion_matrix) && "1" %in% rownames(confusion_matrix)),
        yes = confusion_matrix["1", "1"],
        no = 0 
      ) 
      
      # Verdadeiros negativos
      VN <- ifelse(
        test = ("0" %in% colnames(confusion_matrix) && "0" %in% rownames(confusion_matrix)),
        yes = confusion_matrix["0", "0"],
        no = 0 
      ) 
      
      # Falsos positivos
      FP <- ifelse(
        test = ("1" %in% colnames(confusion_matrix) && "0" %in% rownames(confusion_matrix)),
        yes = confusion_matrix["0", "1"],
        no = 0  
      ) 
      
      # Falsos negativos
      FN <- ifelse(
        test = ("0" %in% colnames(confusion_matrix) && "1" %in% rownames(confusion_matrix)),
        yes = confusion_matrix["1", "0"],
        no = 0
      ) 
      
      # Calcular a acurácia
      sensitivity <- VP/(VP + FN)
      specificity <- VN/(VN + FP)
      # accuracy <- (VP + VN) / (VP + VN + FP + FN)
      accuracy <- sum((predicted_classes == y_train) %>% as.numeric())/ length(y_train) # numero de previsoes corretas / nº total de previsoes
      
      if(is.nan(accuracy) || is.na(accuracy)){
        accuracy <- 0
      }
      
      return(1-accuracy)  # Retornar o negativo da acurácia, pois queremos maximizar
    }
  }
  
  # Carry out grammatical evolution
  ge <- GrammaticalEvolution(
    grammarDef,
    SymRegFitFunc,
    terminationCost = 0.1,
    iterations = 2500,
    max.depth = 5
  )
  
  # Show results
  cat("Grammatical Evolution Search Results:\n")
  cat("  No. Generations: ", ge$generation, "\n")
  cat("  Best Expression: ", deparse(ge$best$expressions[[1]]), "\n")
  cat("  Best Cost: ", ge$best$cost, "\n")
  
  # Evaluate model performance
  env <- as.list(x_test)
  predictions <- try(eval(ge$best$expressions[[1]], envir = env), silent = TRUE)
  
  # Convert continuous forecasts to discrete values ​​(categories)
  predicted_classes <- ifelse(predictions < 0.5, 0, 1)
  # Calculate the confusion matrix
  confusion_matrix <- table(predicted_classes, y_test)
  
  # True positives
  VP <- ifelse(
    test = ("1" %in% colnames(confusion_matrix) && "1" %in% rownames(confusion_matrix)),
    yes = confusion_matrix["1", "1"],
    no = 0 
  ) 
  
  # True negatives
  VN <- ifelse(
    test = ("0" %in% colnames(confusion_matrix) && "0" %in% rownames(confusion_matrix)),
    yes = confusion_matrix["0", "0"],
    no = 0 
  ) 
  
  # False positives
  FP <- ifelse(
    test = ("1" %in% colnames(confusion_matrix) && "0" %in% rownames(confusion_matrix)),
    yes = confusion_matrix["0", "1"],
    no = 0  
  ) 
  
  # False negatives
  FN <- ifelse(
    test = ("0" %in% colnames(confusion_matrix) && "1" %in% rownames(confusion_matrix)),
    yes = confusion_matrix["1", "0"],
    no = 0
  ) 
  
  # Calculate accuracy
  sensitivity <- VP/(VP + FN)
  specificity <- VN/(VN + FP)
  # accuracy <- (VP + VN) / (VP + VN + FP + FN)
  
  accuracy <- sum((predicted_classes == y_test) %>% as.numeric())/ length(y_test) # numero de previsoes corretas / nº total de previsoes
  
  
  cat("Accuracy: ", accuracy, "\n")
  cat("Sensitivity: ", sensitivity, "\n")
  cat("Specificity: ", specificity, "\n")
  
  desempenhos <- tibble(
    BaseDados = nomeBase,
    Fold = fold,
    Acc = accuracy,
    Sens = sensitivity, # com valores NaN
    Spec = specificity, # com valores NaN
  )
  
  resultados$Desempenhos <- resultados$Desempenhos %>% 
    bind_rows(desempenhos)
}

代码验证与问题分析

1. 核心实现的正确性

  • 数据预处理:IRIS数据集的二分类转换逻辑正确,变量命名符合语法规则要求,无错误。
  • 语法规则(ruleDef):动态生成变量规则的逻辑合理,支持二元运算、单变量函数和直接变量引用,符合符号回归的生成逻辑;除法/的潜在错误已被try块捕获,处理正确。
  • SymRegFitFunc适应度函数:
    • 错误处理逻辑覆盖了表达式求值错误、结果类型异常、长度不匹配等情况,返回Inf表示差适应度,逻辑正确。
    • 准确率计算方式直接可靠,且处理了NaN/NA的边界情况;返回1-accuracy符合gramEvol最小化适应度的要求(最大化准确率等价于最小化1-准确率)。

2. 类别不平衡场景的问题根源

代码在IRIS上正常,但在类别不平衡数据集上失效,并非代码实现错误,而是准确率指标不适合类别不平衡场景:

  • 准确率会偏向多数类,模型仅预测多数类就能获得高准确率,导致少数类预测完全失效。
  • 建议替换适应度指标:使用F1-score、AUC,或代价敏感指标(如给FN更高权重);同时固定0.5阈值不合理,可基于训练集ROC曲线动态调整阈值。

3. 潜在优化点

  • 变量索引可改为x_train <- dadosTreinoCV[, -1],避免因数据集列数变化导致的索引错误。
  • 适应度函数中需增加any(is.na(result))的判断,捕获sqrt/log等函数产生的NaN值,避免后续分类错误:
    if (inherits(result, "try-error") || is.language(result) || length(result) != length(y_train) || any(is.na(result))) {
      return(Inf)
    }
    

内容的提问来源于stack exchange,提问作者Juliana Abreu Fontes

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 04:00:53