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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 14:40:15