如何在mlr3 benchmark对象中访问插补后数据以对比插补策略
在mlr3 Benchmark中访问插补后的数据
问题原因
直接访问pipeops$xxx$.result返回NULL,是因为mlr3pipelines在训练完成后不会默认保留中间处理结果——尤其是在交叉验证的benchmark流程中,中间数据不会被持久化存储,仅用于模型训练环节。
解决方案
方法1:从训练好的插补管道重新生成数据
每个训练完成的learner包含了拟合好的插补器,可直接用它处理对应fold的训练数据,得到插补后的结果:
# 取出benchmark中第一个fold的训练好的learner learner = bmr$learners$learner[[1]] # 获取该fold对应的训练数据集 train_task = bmr$resample_results$resample_result[[1]]$task$clone() train_task$filter(bmr$resample_results$resample_result[[1]]$train_set(1)) # 用训练好的插补管道处理数据,得到插补后的任务对象 imputed_task = learner$pipeops$imputeoor$predict(list(train_task))[[1]] # 查看插补后的数据 head(imputed_task$data())
方法2:在管道中添加Debug节点捕获中间数据
如果想在训练过程中直接保存插补结果,可在管道中插入po("debug")节点,强制存储中间数据:
# 修改插补管道,添加debug节点捕获插补后的数据 impute_hist = list( po("missind", type = "integer", affect_columns = selector_type("integer")), po("imputehist", affect_columns = selector_type("integer")) ) %>>% po("featureunion") %>>% po("imputeoor", affect_columns = selector_type("factor")) %>>% po("debug", id = "capture_imputed", keep_results = TRUE) # 新增debug节点 glrn_rf_impute_hist = as_learner(impute_hist %>>% lrn("regr.ranger")) glrn_rf_impute_hist$id = "RF_imp_Hist" # 重新运行benchmark bmr = benchmark(design, store_models = TRUE, store_backends = TRUE) # 从debug节点提取插补后的数据 imputed_data = bmr$learners$learner[[1]]$pipeops$capture_imputed$.result[[1]]$data() head(imputed_data)
方法3:批量提取所有Fold的插补数据
若需要收集交叉验证所有fold的插补结果,可在benchmark时通过extract函数批量捕获:
# 定义提取函数,获取单个resample结果中所有fold的插补数据 extract_imputed = function(resample_result) { lapply(seq_len(resample_result$n_folds), function(i) { learner = resample_result$learners[[i]] train_task = resample_result$task$clone() train_task$filter(resample_result$train_set(i)) imputed_task = learner$pipeops$imputeoor$predict(list(train_task))[[1]] imputed_task$data() }) } # 运行benchmark并指定提取函数 bmr = benchmark(design, store_models = TRUE, store_backends = TRUE, extract = extract_imputed) # 获取所有fold的插补数据 all_imputed_data = bmr$extract_results()[[1]]
对比多种插补策略
只需定义多个插补管道并加入benchmark的learner列表,再用上述方法分别提取数据对比即可:
# 定义第二种插补策略(均值插补) impute_mean = list( po("missind", type = "integer", affect_columns = selector_type("integer")), po("imputemean", affect_columns = selector_type("integer")) ) %>>% po("featureunion") %>>% po("imputeoor", affect_columns = selector_type("factor")) glrn_rf_impute_mean = as_learner(impute_mean %>>% lrn("regr.ranger")) glrn_rf_impute_mean$id = "RF_imp_Mean" # 更新benchmark设计 design = benchmark_grid(tsk_ames, c(glrn_rf_impute_hist, glrn_rf_impute_mean), rsmp_cv3) # 运行benchmark并提取数据 bmr = benchmark(design, store_models = TRUE, store_backends = TRUE, extract = extract_imputed) # 分别获取两种策略的插补数据 hist_imputed = bmr$extract_results()[[1]] mean_imputed = bmr$extract_results()[[2]] # 对比某列的插补结果差异 hist_col = hist_imputed[[1]]$Lot_Frontage mean_col = mean_imputed[[1]]$Lot_Frontage table(hist_col != mean_col)
内容的提问来源于stack exchange,提问作者A. Suliman
相关产品推荐
相关产品推荐

