R语言如何批量为嵌套分组数据创建训练集与测试集?
实现方法
全程不需要解开已经做好的嵌套数据结构,基于purrr::map做分组映射即可批量完成所有子集的训练/测试集拆分,和批量训练线性回归的逻辑完全一致。
推荐写法(适配tidyverse生态)
用rsample包(tidymodels生态中专职做数据拆分的工具包)写拆分逻辑,代码可读性和可扩展性最好,后续做交叉验证、超参数调优都可以无缝衔接:
library(tidyverse) library(datasets) library(rsample) data("ChickWeight") # 原有嵌套逻辑 ChickWeightNest <- ChickWeight %>% group_by(Chick) %>% nest() # 批量拆分训练/测试集 ChickWeightSplit <- ChickWeightNest %>% mutate( # 按7:3比例拆分单只雏鸡的数据,需要分层抽样可给strata参数传入对应列名 split = map(data, ~initial_split(.x, prop = 0.7)), train_set = map(split, training), test_set = map(split, testing) )
拆分完成后直接继续用mutate+map即可批量完成模型训练、预测、评估全流程,替换成xgboost模型也不需要改整体结构,示例:
ChickWeightRes <- ChickWeightSplit %>% mutate( # 这里替换成xgboost的训练逻辑即可,输入用train_set列 model = map(train_set, ~lm(weight ~ Time, data = .x)), # 测试集预测 pred = map2(model, test_set, predict), # 计算测试集RMSE作为评估指标 rmse = map2_dbl(pred, test_set, ~sqrt(mean((.x - .y$weight)^2))) )
无额外依赖的基础写法
如果不想装其他包,用基础R的抽样逻辑也能实现同样效果:
ChickWeightSplitBase <- ChickWeightNest %>% mutate( train_index = map(data, ~sample(seq_len(nrow(.x)), size = round(nrow(.x)*0.7))), train_set = map2(data, train_index, ~.x[.y, ]), test_set = map2(data, train_index, ~.x[-.y, ]) )
注意事项
ChickWeight数据集单只雏鸡的样本量非常少(单组最多12条观测),用xgboost这类树模型时注意控制树深、正则项参数,避免严重过拟合- 如果后续需要做交叉验证、超参数网格调优,只需要把
initial_split替换成vfold_cv、bootstraps等重抽样函数,整体嵌套逻辑不需要改动
学习参考
- 核心知识点是tidyverse的列表列和purrr映射逻辑,所有分组级别的操作都可以封装为自定义函数传入map,不需要写显式for循环
- tidymodels官方教程的「Many Models」章节专门覆盖这类分组批量建模场景,包含从数据嵌套、拆分到训练评估的完整流程
- 检索相关内容可以用关键词
R nested tibble train test split purrr、R batch train xgboost on nested data,能找到大量可复现的实践案例
内容的提问来源于stack exchange,提问作者887
相关产品推荐
相关产品推荐

