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
相关产品推荐
相关产品推荐

