使用caret::train结合bnlearn自定义模型时出现invalid node(s) 'xNames'错误
解决caret+bnlearn自定义模型训练中的"invalid node(s) 'xNames'"错误
问题重现
使用caret::train()结合bnlearn训练自定义贝叶斯网络模型,网格搜索不同bnlearn算法时出现如下错误:
"model fit failed for Fold1: algorithm=hc Error in check.nodes(name, x) : invalid node(s) 'xNames'. " Warning message: "model fit failed for Fold1: algorithm=tabu Error in check.nodes(name, x) : invalid node(s) 'xNames'. "
数据集可通过以下代码获取:
install.packages("mlr3fairness") library("mlr3") data("adult_test", package = "mlr3fairness") adult_test <- adult_test %>% rename(income = target)
核心调用代码:
model <- caret::train(income ~ ., data = outer_trainData, method = bn_model, tuneGrid = tunegrid, trControl = inner_control)
错误原因
错误根源在自定义模型bn_model的predictors函数:
predictors = function(x, ...) { colnames(x) }
当caret执行交叉验证时,会调用这个函数提取模型的特征变量名,但这里的x是bnlearn训练出的bn.fit对象,而非原始数据框。colnames(x)对bn.fit对象会返回'xNames'这个无效节点名,导致后续bnlearn的节点检查失败。
解决方案
修改bn_model中的predictors函数,改为从bn.fit对象中提取正确的节点名称:
predictors = function(x, ...) { # 从bn.fit对象中获取所有节点名 nodes(x) }
同时建议在predict函数中增加列名一致性检查,避免新数据列名与模型节点不匹配:
predict = function(modelFit, newdata, preProc = NULL, submodels = NULL) { if (is.null(modelFit)) { return(rep(NA, nrow(newdata))) } data <- discretize_df(newdata) # 确保新数据列名和模型节点一致 if (!all(colnames(data) %in% nodes(modelFit))) { stop("New data columns do not match model nodes.") } predictions <- tryCatch({ predict(modelFit, data) }, error = function(e) { print(paste("Error in prediction: ", e$message)) return(rep(NA, nrow(newdata))) }) return(predictions) }
完整修正后的bn_model片段
bn_model <- list( label = "Bayesian Network", library = "bnlearn", type = "Classification", parameters = data.frame( parameter = c("algorithm"), class = c("character"), label = c("Algorithm") ), grid = function(x, y, len = NULL, search = "grid") { algorithms <- c("hc", "tabu", "gs", "iamb") if (search == "grid") { expand.grid(algorithm = algorithms) } else { data.frame(algorithm = sample(algorithms, len, replace = TRUE)) } }, fit = function(x, y, wts, param, lev, last, classProbs, ...) { df <- as.data.frame(x) df$income <- y # Ensure consistent factor levels across folds df <- lapply(df, function(col) { if (is.factor(col)) { levels(col) <- union(levels(col), unique(col)) } return(col) }) df <- as.data.frame(df) tryCatch({ # Train the Bayesian network using the specified algorithm bn <- train_bn(df, param$algorithm) # Fit the parameters of the Bayesian network bn_fitted <- bn.fit(bn, df) return(bn_fitted) }, error = function(e) { print(paste("Error in fitting model:", e$message)) return(NULL) }) }, predict = function(modelFit, newdata, preProc = NULL, submodels = NULL) { if (is.null(modelFit)) { return(rep(NA, nrow(newdata))) } data <- discretize_df(newdata) # 确保新数据列名和模型节点一致 if (!all(colnames(data) %in% nodes(modelFit))) { stop("New data columns do not match model nodes.") } predictions <- tryCatch({ predict(modelFit, data) }, error = function(e) { print(paste("Error in prediction: ", e$message)) return(rep(NA, nrow(newdata))) }) return(predictions) }, prob = NULL, predictors = function(x, ...) { nodes(x) }, sort = function(x) x, levels = function(x) levels(x$obs) )
内容的提问来源于stack exchange,提问作者foesine
相关产品推荐
相关产品推荐

