含多水平分类特征时,如何在R中划分训练集与测试集
在R中确保分类变量全水平存在的训练/测试集划分方法
当分类变量存在多水平且无法合并时,直接随机划分可能导致测试集出现训练集未见过的水平,引发预测错误。以下是几种可靠的解决方法:
基础R手动实现
适用于需要精细控制的场景,核心思路是先保证每个分类水平至少有1个样本进入训练集,再对剩余样本按比例抽取:
# 假设数据集为df,目标分类变量为cat_var # 按分类水平分组提取索引 level_groups <- split(seq_len(nrow(df)), df$cat_var) # 每个水平先选1个样本进入训练集 train_base <- unlist(lapply(level_groups, function(idx) sample(idx, size = 1))) # 剩余待划分样本的索引 remaining_idx <- setdiff(seq_len(nrow(df)), train_base) # 按目标比例(比如70%训练集)抽取剩余训练样本 train_add <- sample(remaining_idx, size = round(0.7 * nrow(df)) - length(train_base)) # 合并得到最终训练集索引 train_idx <- c(train_base, train_add) train_df <- df[train_idx, ] test_df <- df[-train_idx, ] # 验证所有水平在两个集合中都存在 all(levels(df$cat_var) %in% levels(train_df$cat_var)) all(levels(df$cat_var) %in% levels(test_df$cat_var))
使用caret包分层抽样
caret是机器学习常用工具包,createDataPartition函数支持按指定变量分层,确保每个水平都能分配到训练/测试集:
library(caret) # 按cat_var分层,训练集比例设为0.7 train_idx <- createDataPartition(df$cat_var, p = 0.7, list = FALSE) train_df <- df[train_idx, ] test_df <- df[-train_idx, ] # 验证水平完整性 table(train_df$cat_var) table(test_df$cat_var)
如果存在多个分类变量,可通过interaction()组合变量,确保各变量水平的交叉组合都能覆盖:
# 按cat_var1和cat_var2的组合分层 train_idx <- createDataPartition(interaction(df$cat_var1, df$cat_var2), p = 0.7, list = FALSE)
使用rsample包(tidyverse风格)
rsample提供更简洁的tidy风格接口,initial_split的strata参数可指定分层变量:
library(rsample) # 按cat_var分层拆分数据集,训练集比例0.7 data_split <- initial_split(df, prop = 0.7, strata = cat_var) train_df <- training(data_split) test_df <- testing(data_split) # 验证水平存在性 all(levels(df$cat_var) %in% levels(test_df$cat_var))
注意事项
- 若某分类水平仅含1个样本,需直接将其划入训练集,否则测试集出现该水平时模型无法处理;
- 划分完成后务必验证两个集合的分类水平与原数据集一致,避免遗漏。
内容的提问来源于stack exchange,提问作者user2899944
相关产品推荐
相关产品推荐

