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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 06:50:55