使用R的tidymodels训练BART模型后对象异常问题求助
使用tidymodels训练BART模型的异常问题
我用tidymodels框架训练BART模型时遇到两个异常:
- 原本正常的模型对象
bart_mod在拟合工作流后变成"call: NULL",我没直接修改过它; - 拟合后的模型
bart_fit无法查看任何信息,也没有对应的tidy方法,但却能正常预测。
以下是可复现代码:
library(tidyverse) library(tidymodels) set.seed(2022) # Parameters -------------------------------------------------------------- n <- 5000 coef_x_var_1 <- 1 coef_x_var_2 <- 2 coef_x_var_3 <- 3 gen_y_1 <- function(data = dataset) { return(data$y_0 + data$x_var_1*coef_x_var_1 + data$x_var_2*coef_x_var_2 + data$x_var_3*coef_x_var_3 + rnorm(n = nrow(data), mean = 0, sd = 3) )} # Data generation --------------------------------------------------------- dataset <- matrix(NA, nrow = n, ncol = 3) # Generate the unit-level moderators dataset[,1] <- rnorm(mean = rnorm(n = 1), n = n) dataset[,2] <- rnorm(mean = rnorm(n = 1), n = n) dataset[,3] <- rnorm(mean = rnorm(n = 1), n = n) # Change into dataframe colnames(dataset) <- c("x_var_1", "x_var_2", "x_var_3") dataset <- as_tibble(dataset) # Make sure the variable format is numeric (except for the identifiers) dataset$x_var_1 <- as.numeric(dataset$x_var_1) dataset$x_var_2 <- as.numeric(dataset$x_var_2) dataset$x_var_3 <- as.numeric(dataset$x_var_3) # Generate the untreated potential outcomes P0_coefs <- rdunif(n = 6, 1, 15) dataset$y_0 <- dataset$x_var_1*P0_coefs[4] + dataset$x_var_2*P0_coefs[5] + dataset$x_var_3*P0_coefs[6] + rnorm(n = nrow(dataset), mean = 0, sd = 3) dataset$y_1 <- gen_y_1(data = dataset) # Create a variable to indicate treatment treatment_group <- sample(1:nrow(dataset), size = nrow(dataset)/2) # Indicate which potential outcome you observe obs_dataset <- dataset |> mutate(treated = ifelse(row_number() %in% treatment_group, 1, 0), obs_y = ifelse(treated, y_1, y_0)) y1_obs_dataset <- obs_dataset |> filter(treated == 1) y0_obs_dataset <- obs_dataset |> filter(treated == 0) # Analysis ---------------------------------------------------------------- covariates <- c("x_var_1", "x_var_2", "x_var_3") bart_formula <- as.formula(paste0("obs_y ~ ", paste(covariates, collapse = " + "))) # Create the workflow bart_mod <- bart() |> set_engine("dbarts") |> set_mode("regression") bart_recipe <- recipe(bart_formula, data = obs_dataset) |> step_zv(all_predictors()) bart_workflow <- workflow() |> add_model(bart_mod) |> add_recipe(bart_recipe) # The workflow first looks right bart_workflow #> ══ Workflow ════════════════════════════════════════════════════════════════════ #> Preprocessor: Recipe #> Model: bart() #> #> ── Preprocessor ──────────────────────────────────────────────────────────────── #> 1 Recipe Step #> #> • step_zv() #> #> ── Model ─────────────────────────────────────────────────────────────────────── #> BART Model Specification (regression) #> #> Computational engine: dbarts # Once I fit it though, the model part becomes call: NULL bart_fit <- bart_workflow |> fit(y1_obs_dataset) # Nothing is stored in the fit bart_fit #> ══ Workflow [trained] ══════════════════════════════════════════════════════════ #> Preprocessor: Recipe #> Model: bart() #> #> ── Preprocessor ──────────────────────────────────────────────────────────────── #> 1 Recipe Step #> #> • step_zv() #> #> ── Model ─────────────────────────────────────────────────────────────────────── #> #> Call: #> `NULL`() # The content of this object has changed! bart_workflow #> ══ Workflow ════════════════════════════════════════════════════════════════════ #> Preprocessor: Recipe #> Model: bart() #> #> ── Preprocessor ──────────────────────────────────────────────────────────────── #> 1 Recipe Step #> #> • step_zv() #> #> ── Model ─────────────────────────────────────────────────────────────────────── #> #> Call: #> NULL bart_fit |> extract_fit_parsnip(bart_fit) #> parsnip model object #> #> #> Call: #> `NULL`() # And yet, I am able to run a prediction using the fit! predict(bart_fit, y0_obs_dataset) #> # A tibble: 2,500 × 1 #> .pred #> <dbl> #> 1 -4.67 #> 2 -6.23 #> 3 6.35 #> 4 10.7 #> 5 4.90 #> 6 -13.8 #> 7 4.70 #> 8 19.6 #> 9 -0.907 #> 10 5.38 #> # … with 2,490 more rows
Created on 2022-12-24 with reprex v2.0.2
问题原因及解决方法
1. 原始模型对象bart_mod变为call: NULL的原因
工作流默认会引用并修改传入的模型对象,而非创建副本。调用fit()时,工作流会修改bart_mod的内部状态,导致其显示异常。
解决方法:将模型副本传入工作流,避免原对象被修改:
# 方法一:创建独立副本 bart_mod_copy <- bart() |> set_engine("dbarts") |> set_mode("regression") bart_workflow <- workflow() |> add_model(bart_mod_copy) |> add_recipe(bart_recipe) # 方法二:使用clone()复制对象 bart_workflow <- workflow() |> add_model(clone(bart_mod)) |> add_recipe(bart_recipe)
2. 无法查看拟合后模型信息的解决方法
拟合后的核心模型结果存储在工作流的底层引擎对象中,需要用特定函数提取:
- 使用
extract_fit_engine()提取真正的dbarts拟合结果:
# 提取底层dbarts模型 bart_engine_fit <- extract_fit_engine(bart_fit) # 查看模型基本信息 print(bart_engine_fit) # 获取变量重要性 variable_importance <- bartEngine::varImp(bart_engine_fit) print(variable_importance) # 查看拟合细节统计 summary(bart_engine_fit)
- 关于tidy方法:目前parsnip对dbarts的tidy支持有限,可尝试用
broom.mixed包整理结果:
library(broom.mixed) tidy(bart_engine_fit)
补充说明
tidymodels的设计逻辑是封装训练流程,所以拟合后的工作流会隐藏底层细节,但核心拟合结果并未丢失,只需用对应函数提取。避免原模型对象被修改的关键是传递副本而非原对象引用。
内容的提问来源于stack exchange,提问作者Martin
相关产品推荐
相关产品推荐

