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

tidymodels+bonsai构建ctree后,ggparty无法读取模型species数据的问题

解决tidymodels+bonsai构建ctree后ggparty绘图变量找不到的问题

问题原因

tidymodels工作流默认会剥离训练数据以节省内存,通过extract_fit_parsnip()提取的封装模型不包含原始拟合数据;而直接用partykit构建的模型会默认保留数据,这才是ggparty的geom_node_plot能直接调用变量的原因。

解决方案

要提取包含拟合数据的原生partykit模型,需分两步操作:

1. 拟合时保留训练数据

在最终拟合工作流时,通过control_parsnip(keep_data = TRUE)参数强制保留训练数据。示例代码:

# 加载所需包
library(tidymodels)
library(bonsai)
library(ggparty)

# 数据拆分(以iris为例)
data_split <- initial_split(iris, prop = 0.7)
train_data <- training(data_split)
test_data <- testing(data_split)

# 定义ctree模型并设置调优参数
ctree_spec <- decision_tree(tree_depth = tune()) %>%
  set_engine("partykit") %>%
  set_mode("classification")

# 构建工作流
wf <- workflow() %>%
  add_formula(Species ~ .) %>%
  add_model(ctree_spec)

# 参数调优(示例用网格搜索)
grid <- grid_regular(tree_depth(range = c(1, 3)), levels = 3)
tune_res <- tune_grid(wf, resamples = vfold_cv(train_data), grid = grid)
best_params <- select_best(tune_res, metric = "accuracy")

# 最终拟合,设置保留数据
fit_final <- wf %>%
  finalize_workflow(best_params) %>%
  fit(data = train_data, control = control_parsnip(keep_data = TRUE))

2. 提取原生partykit模型

不要用extract_fit_parsnip(),改用extract_fit_engine()提取底层的partykit::ctree对象,这个对象包含训练数据:

ctree_party <- extract_fit_engine(fit_final)

3. 用ggparty绘图

现在可以正常使用geom_node_plot绘制终端节点的物种条形图:

ggparty(ctree_party) +
  geom_node_splitvar() +
  geom_node_plot(
    gglist = list(
      geom_bar(aes(x = Species, fill = Species), stat = "count"),
      theme_minimal()
    ),
    shared_axis_labels = TRUE
  )

补充说明

如果使用last_fit()做最终拟合,需在control_last_fit中传递control_parsnip参数:

last_fit_res <- wf %>%
  finalize_workflow(best_params) %>%
  last_fit(split = data_split, control = control_last_fit(control_parsnip = control_parsnip(keep_data = TRUE)))

ctree_party <- last_fit_res %>% extract_fit_engine()

内容的提问来源于stack exchange,提问作者Jacek Kotowski

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 11:36:28