使用tidymodels计算全类别预测值 批量生成分组预测列
tidymodels 批量生成多group预测值实现方案
依赖包加载
library(tidyverse) library(tidymodels)
示例数据构造(与原参考逻辑一致)
year <- rep(2014:2015, length.out=10000) group <- sample(c(0,1,2,3,4,5,6), replace=TRUE, size=10000) value <- sample(10000, replace=T) female <- sample(c(0,1), replace=TRUE, size=10000) smoker <- sample(c(0,1), replace=TRUE, size=10000) dta <- data.frame(year=year, group=group, value=value, female=female, smoker=smoker)
分组拟合模型(适配原参考代码的按year+group拆分建模逻辑)
# 构造嵌套分组数据集 dta_nested <- dta %>% group_by(year, group) %>% nest() %>% # 定义probit二分类模型并批量拟合 mutate( model = map(data, ~{ glm_spec <- logistic_reg(mode = "classification") %>% set_engine("glm", family = binomial(link = "probit")) fit(glm_spec, smoker ~ female*group, data = .x) }) )
批量生成7组预测结果
pred_results <- map_dfc(0:6, function(target_group) { # 批量替换group为当前遍历值,无需手动构造7个独立数据集 new_data <- dta %>% mutate(group = target_group) # 匹配对应分组模型完成预测 pred_vec <- new_data %>% left_join(dta_nested, by = c("year", "group")) %>% mutate(pred = map2_dbl(model, data.x, ~predict(.x, new_data = .y, type = "prob")$.pred_1)) %>% pull(pred) # 返回单列预测结果,列名对应group取值 tibble(!!paste0("pred_group_", target_group) := pred_vec) }) # 合并预测结果到原数据集,直接得到带7列预测值的完整数据 dta_with_pred <- bind_cols(dta, pred_results)
逻辑说明
- 无需手动构造
dat_new0到dat_new67个独立数据集,通过purrr::map遍历group取值自动完成数据集替换 - 预测逻辑批量执行,最终结果自动按列合并,如需调整group取值范围,仅需修改
0:6为对应取值序列即可 - 完全适配原参考代码的按year+group分组建模逻辑,预测时自动匹配对应子模型保证结果一致性
内容的提问来源于stack exchange,提问作者Stata_user
相关产品推荐
相关产品推荐

