如何利用mlr3中resample存储的模型对新数据生成样本外预测?
解决mlr3交叉拟合后复用模型预测新数据的问题
报错原因
Error: No task stored, and no task provided的核心原因是:通过resample生成的训练后学习器默认不绑定原任务,而mlr3的predict()方法需要明确的任务对象支撑,无法直接接收原始数据框作为输入。
两种可行解决方案
方案1:将新数据转为Task后传入预测
把修改后的dat_1转换成对应的回归任务,再传入predict()方法:
# 创建新数据对应的回归任务 task_1 <- as_task_regr(dat_1, target = "Y") # 对第一折的测试集进行预测 predict(rr$learners[[1]], task = task_1, row_ids = rr$resampling$test_set(1))
方案2:使用predict_newdata()直接传入数据框
mlr3提供了predict_newdata()函数,可直接接收数据框,无需手动创建Task:
# 对第一折测试集的新数据进行预测 predict_newdata(rr$learners[[1]], newdata = dat_1[rr$resampling$test_set(1), ])
完整交叉拟合预测流程示例
如果要批量完成K折的交叉拟合预测(对dat_1的每个测试集用对应模型预测),可以用循环实现:
# 初始化存储预测结果的列表 pred_list <- vector("list", K) # 遍历每折 for (k in 1:K) { # 获取当前折的测试集行ID test_ids <- rr$resampling$test_set(k) # 用对应模型预测新数据的测试集部分 pred_list[[k]] <- predict_newdata(rr$learners[[k]], newdata = dat_1[test_ids, ]) } # 合并所有预测结果 all_preds <- do.call(rbind, pred_list)
额外优化:让学习器自动存储任务(可选)
如果希望训练后的学习器自动绑定任务,可以在创建学习器时设置store_task = TRUE,后续调用predict()时无需额外传入Task:
# 创建学习器时开启任务存储 learn_gbm <- lrn("regr.lightgbm", store_task = TRUE) # 重新执行resample rr <- resample(task, learn_gbm, cv, store_models = TRUE) # 直接用predict完成预测 predict(rr$learners[[1]], row_ids = rr$resampling$test_set(1), newdata = dat_1[rr$resampling$test_set(1), ])
内容的提问来源于stack exchange,提问作者N. Williams
相关产品推荐
相关产品推荐

