关于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
相关产品推荐
相关产品推荐

