如何在tidymodels中使用自定义拆分数据结合fit_resamples?
自定义交叉验证集在tidymodels中的正确使用方法
问题描述
我有一个自定义函数,能依据多种规则拆分训练集和测试集,想在tidymodels工作流里结合fit_resamples使用,但模仿vfold_cv的结构后还是报错,提示需要rset对象,强制设类也没用。
用户示例代码
data(ames, package = "modeldata") split_data <- function(df, n) { set.seed(123) # for reproducibility df$id <- seq.int(nrow(df)) list_of_splits <- list() for(i in 1:n) { train_index <- sample(df$id, size=ceiling(nrow(df)*.8)) train_set <- df[train_index,] test_set <- df[-train_index,] list_of_splits[[i]] <- list(train_set = train_set, test_set = test_set) } return(list_of_splits) } splits <- split_data(ames, 5) resamples <- map(splits, ~rsample::make_splits( x = .$train_set |> select(colnames(.$test_set)), assessment = .$test_set )) names(resamples) <- paste0("Fold", seq_along(resamples)) resamples <- tibble::tibble(splits = resamples, id = names(resamples)) lm_model <- linear_reg() %>% set_engine("lm") lm_wflow <- workflow() %>% add_model(lm_model) %>% add_formula(Sale_Price ~ Longitude + Latitude) res <- lm_wflow %>% fit_resamples(resamples = resamples)
报错信息
Error in `check_rset()`: ! The `resamples` argument should be an 'rset' object, such as the type produced by `vfold_cv()` or other 'rsample' functions.
解决方案
你需要用rsample::manual_rset()来把自定义的splits转换成合法的rset对象,而不是直接构建tibble。具体修改如下:
library(tidyverse) library(tidymodels) data(ames, package = "modeldata") split_data <- function(df, n) { set.seed(123) # 保证可复现 df$id <- seq.int(nrow(df)) list_of_splits <- list() for(i in 1:n) { train_index <- sample(df$id, size=ceiling(nrow(df)*.8)) train_set <- df[train_index,] test_set <- df[-train_index,] # 直接生成rsplit对象,省去后续转换步骤 list_of_splits[[i]] <- make_splits( x = train_set %>% select(-id), # 移除手动添加的id列,避免干扰模型 assessment = test_set %>% select(-id) ) } return(list_of_splits) } # 生成拆分后的rsplit对象列表 splits_list <- split_data(ames, 5) # 用manual_rset包装成rset对象,这一步是关键 custom_resamples <- manual_rset( splits = splits_list, ids = paste0("Fold", seq_along(splits_list)) ) # 原有工作流无需修改 lm_model <- linear_reg() %>% set_engine("lm") lm_wflow <- workflow() %>% add_model(lm_model) %>% add_formula(Sale_Price ~ Longitude + Latitude) # 现在可以正常运行fit_resamples res <- lm_wflow %>% fit_resamples(resamples = custom_resamples) # 查看评估结果 collect_metrics(res)
关键要点:
manual_rset()是rsample包专门用于创建自定义交叉验证集的工具,会自动赋予对象正确的rset类属性,让tidymodels可以识别- 记得移除手动添加的
id列,避免它被当作特征传入模型
额外问题解答:各折训练/测试集大小略有差异是否有影响?
这种小差异一般不会有严重影响:
- 只要差异不是极端悬殊(比如某折训练集仅占总数据的50%,另一折占90%),模型评估结果的稳定性不会受太大干扰
- 交叉验证的核心是用不同拆分验证泛化能力,少量样本量波动属于可接受范围
- 如果差异是业务规则导致的(比如按分组拆分,每组大小不同),只要拆分逻辑符合需求,无需刻意调整
如果担心影响评估结果,可以:
- 拆分时尽量保证各折样本量接近(比如用分层抽样,或调整抽样比例)
- 增加交叉验证的折数,用更多拆分来稀释单折样本量差异的影响
内容的提问来源于stack exchange,提问作者acircleda
相关产品推荐
相关产品推荐

