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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 17:43:12