如何在R语言中创建10折交叉验证时确保相同ID的样本处于同一折?
解决同一ID样本归为同一折的10折交叉验证问题
这个问题很常见——当数据存在重复ID(比如同一主体的多个观测)时,直接按行拆分交叉验证会破坏数据的关联性。下面是具体的解决方法,用caret包就能轻松实现你的需求:
核心思路
先对唯一ID集合进行10折划分,再将每个ID对应的所有样本映射到对应的折中,这样就能保证同一个ID的所有样本都在同一折里。
具体实现步骤
提取唯一ID列表
先从数据框中取出所有不重复的ID,我们要对这些ID做折划分:unique_ids <- unique(df$ID)对唯一ID创建交叉验证折
使用caret包的createFolds()函数,这次针对唯一ID而不是原始数据的行:library(caret) # 设置随机种子保证结果可复现 set.seed(123) id_folds <- createFolds(unique_ids, k = 10, list = TRUE)将ID折映射为原始数据的行索引
遍历每个折的ID,找到原始数据中对应这些ID的所有行,生成最终的折索引:fold_indices <- lapply(id_folds, function(fold_ids) { which(df$ID %in% unique_ids[fold_ids]) })提取训练集和测试集
现在你就可以用fold_indices来拆分数据了,比如取第一折作为测试集:test.data <- df[fold_indices[[1]], ] train.data <- df[-fold_indices[[1]], ]
验证结果
你可以用下面的代码验证,确保测试集和训练集的ID没有重叠:
# 检查交集长度是否为0,0则说明无重叠 length(intersect(test.data$ID, train.data$ID)) == 0
不用caret包的base R实现
如果你不想依赖caret包,也可以用base R手动实现:
set.seed(123) unique_ids <- unique(df$ID) # 给唯一ID随机分组为10折 id_groups <- split(unique_ids, sample(rep(1:10, length.out = length(unique_ids)))) # 转换为行索引 fold_indices <- lapply(id_groups, function(ids) which(df$ID %in% ids))
内容的提问来源于stack exchange,提问作者Babak Kasraei
相关产品推荐
相关产品推荐

