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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.27 05:24:02