如何用R的stacks包堆叠不同数据集模型并解决重采样错误?
不同数据集模型堆叠的实现方案
我正在使用两个不同国家的数据集,目标是集成两国的模型以观察集成模型的泛化能力。当前设置为:为每个国家训练了一个workflow_set(包含10个模型规格,带有重采样及规模为20的网格搜索)。尝试通过以下代码将它们作为候选模型添加时出现错误:
predictions <- stacks() %>% add_candidates(wf_set_1) %>% add_candidates(wf_set_2)
错误信息:
Error:
It seems like the new candidate member 'Logistic Regression' doesn't make use of the same resampling object as the existing candidates.
错误原因与解决方法
错误原因
stacks包的add_candidates要求所有候选模型必须基于完全相同的重采样折分对象,因为堆叠模型需要依赖同一组重采样的预测结果来训练元学习器。你当前的两个workflow_set分别基于各自数据集的重采样,不符合这个要求,因此报错。
可行方案
方案一:合并数据集并使用统一重采样
这是最贴合stacks包设计逻辑的方案:
- 将两个国家的数据集合并,新增一个标识样本所属国家的特征
- 基于合并后的数据集创建统一的重采样对象(比如交叉验证折分)
- 基于这个统一重采样重新训练两个国家的
workflow_set - 之后即可正常使用
add_candidates完成堆叠
示例代码:
# 合并数据集,添加国家标识 library(dplyr) combined_data <- bind_rows( country1_data %>% mutate(country = "country1"), country2_data %>% mutate(country = "country2") ) # 创建统一的10折交叉验证重采样 set.seed(123) common_resample <- rsample::vfold_cv(combined_data, v = 10) # 重新训练workflow_set,使用统一重采样 wf_set_1 <- workflow_set(...) %>% workflowsets::workflow_map( "tune_grid", resamples = common_resample, grid = 20 ) wf_set_2 <- workflow_set(...) %>% workflowsets::workflow_map( "tune_grid", resamples = common_resample, grid = 20 ) # 正常执行堆叠操作 predictions <- stacks::stacks() %>% stacks::add_candidates(wf_set_1) %>% stacks::add_candidates(wf_set_2)
方案二:手动构建堆叠元模型(绕过stacks包限制)
如果不想合并数据集,可手动完成堆叠的核心逻辑:
- 在各自数据集上训练好所有模型后,用一个独立的公共验证集(可以是合并后的测试集,或者其中一个国家的数据集)收集所有模型的预测结果
- 将这些预测结果作为元特征,结合真实标签训练元模型(比如逻辑回归、随机森林),以此整合不同模型的输出
示例代码:
# 假设已训练好所有模型:model1_1~model1_10(来自国家1)、model2_1~model2_10(来自国家2) # 准备公共验证集(这里用合并数据集的20%作为验证) validation_data <- combined_data %>% dplyr::slice_sample(prop = 0.2) # 收集所有模型的预测结果作为元特征 meta_features <- validation_data %>% mutate( # 国家1模型的概率预测 pred_m1_1 = predict(model1_1, ., type = "prob")$.pred_positive, pred_m1_2 = predict(model1_2, ., type = "prob")$.pred_positive, # ... 补充其他国家1的模型 # 国家2模型的概率预测 pred_m2_1 = predict(model2_1, ., type = "prob")$.pred_positive, pred_m2_2 = predict(model2_2, ., type = "prob")$.pred_positive # ... 补充其他国家2的模型 ) %>% select(starts_with("pred_"), true_label) # 保留元特征和真实标签 # 训练元模型(以逻辑回归为例) meta_model <- parsnip::logistic_reg() %>% parsnip::set_engine("glm") %>% parsnip::fit(true_label ~ ., data = meta_features) # 使用元模型生成最终预测 final_predictions <- predict(meta_model, meta_features, type = "class")
内容的提问来源于stack exchange,提问作者Joao Souza
相关产品推荐
相关产品推荐

