dplyr中case_when调用n()配合slice_sample抽样报错解决
报错原因
报错由dplyr的参数规则和函数调用环境限制导致:
slice_sample()的n参数仅接受长度为1的常量值,不支持传入随分组动态变化的向量,直接在该参数位置写case_when动态逻辑不符合参数输入要求n()是dplyr上下文专用函数,仅能在mutate、summarise、filter等支持逐组计算的动词内部的数据掩码环境中调用,放在slice_sample()的n参数位置不属于合法调用环境,因此触发对应报错- 你代码中
mutate(counts = n())生成的分组计数列不会被slice_sample()的n参数识别,该参数不会从待处理的数据集中读取取值,因此这一步对后续抽样逻辑没有实际作用
实现方案
方案1:按组调用slice_sample(逻辑最直观)
通过group_modify()逐组执行抽样逻辑,每个分组内可以直接获取组样本量做判断,完全沿用slice_sample()的抽样逻辑:
library(dplyr) dat <- mtcars |> group_by(cyl) |> group_modify(~ { sample_size <- if (nrow(.x) > 10) 10 else nrow(.x) slice_sample(.x, n = sample_size) })
代码中.x代表当前分组对应的子集数据,nrow(.x)即为当前分组的样本量,判断后传入对应抽样数即可,完全匹配“组样本量大于10抽10个,否则抽全部”的需求。
方案2:随机排序后筛选(性能更优,适合大数据集)
不需要嵌套函数,通过给组内样本生成随机序号再筛选的方式实现等效抽样,执行效率更高:
library(dplyr) dat <- mtcars |> group_by(cyl) |> # 为组内所有样本生成不重复的随机排序序号 mutate(random_id = sample(1:n())) |> # 每组保留序号<=10的样本,样本量不足10的组会自动保留全部样本 filter(random_id <= 10) |> select(-random_id)
该方案的抽样效果和slice_sample()默认的无放回简单随机抽样完全一致,大样本量场景下优先选择。
内容的提问来源于stack exchange,提问作者Peter
相关产品推荐
相关产品推荐

