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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 19:33:25