R语言base包sample无放回采样忽略小概率问题排查
R语言
base::sample()无放回采样的小概率截断问题 问题现象
在无放回采样场景中,base::sample()函数存在非零小概率样本被"截断"的问题:部分理论上有被选中概率的样本,在重复大量采样后从未被选中。对比手动实现的逐次采样(每次抽取后移除已选样本并重新归一化权重),内置函数的结果明显不符合预期,且对权重进行10^i倍数缩放后问题依然存在。
实验验证
实验代码
N_simulations = 10000 N_draw_per_sim = 10 dat = data.frame(id = 1:40, log_likelihood = seq(from = 550, to = 350, length.out = 40)) dat$lkhd_wt = (function(x) {x / sum(x)}) (exp(dat$log_likelihood - max(dat$log_likelihood))) sample_base = function(N_) {sample(x = dat$id, size = N_, prob = dat$lkhd_wt, replace = FALSE)} sample_manual = function(N_) { # 手动逐次采样,移除已选样本后重新加权 smpl_ = sample(x = dat$id, size = 1, prob = (function(x) {x / sum(x)}) (dat$lkhd_wt), replace = FALSE) for(i in 2:N_) { prv_idx = which(dat$id %in% smpl_) smpl_ = c(smpl_, sample(x = dat$id[-prv_idx], size = 1, prob = (function(x) {x / sum(x)}) (dat$lkhd_wt[-prv_idx]), replace = FALSE)) } return(smpl_) } # 运行内置函数采样 dat_sim_draw_basic = as.data.frame(matrix(0, nrow = dim(dat)[1], ncol = N_simulations)) for(sim_i in 1:N_simulations) {dat_sim_draw_basic[sample_base(N_draw_per_sim), sim_i] = 1} # 运行手动采样 dat_sim_draw_manual = as.data.frame(matrix(0, nrow = dim(dat)[1], ncol = N_simulations)) for(sim_i in 1:N_simulations) {dat_sim_draw_manual[sample_manual(N_draw_per_sim), sim_i] = 1} # 统计采样结果 dat_agg_basic = data.frame(id = dat$id, cnt = apply(dat_sim_draw_basic, 1, sum)) dat_agg_basic$wt = dat_agg_basic$cnt / sum(dat_agg_basic$cnt) dat_agg_manual = data.frame(id = dat$id, cnt = apply(dat_sim_draw_manual, 1, sum)) dat_agg_manual$wt = dat_agg_manual$cnt / sum(dat_agg_manual$cnt) # 合并对比 TF_non_zero = dat_agg_basic$wt != 0 | dat_agg_manual$wt != 0 dat_compare = merge(x = dat_agg_basic[TF_non_zero,], y = dat_agg_manual[TF_non_zero,], by.x = c("id"), by.y = c("id"), all.x = TRUE, all.y = TRUE) colnames(dat_compare) = c("id", "cnt_basic", "wt_basic", "cnt_manual", "wt_manual") dat_compare = dat_compare[,c("id", "cnt_basic", "cnt_manual", "wt_basic", "wt_manual")] dat_compare
实验输出
id cnt_basic cnt_manual wt_basic wt_manual 1 1 10000 10000 0.1 0.10000 2 2 10000 10000 0.1 0.10000 3 3 10000 10000 0.1 0.10000 4 4 10000 10000 0.1 0.10000 5 5 10000 10000 0.1 0.10000 6 6 10000 10000 0.1 0.10000 7 7 10000 10000 0.1 0.10000 8 8 10000 10000 0.1 0.10000 9 9 10000 10000 0.1 0.10000 10 10 10000 9958 0.1 0.09958 11 11 0 42 0.0 0.00042
可以看到,id=11的样本在base::sample()中从未被选中(cnt_basic=0),但手动采样有42次选中记录。
原因分析
base::sample()在无放回带权重采样时,采用的是指数分布排序法:对每个样本生成-log(runif(n)) / prob的值,排序后取前size个样本。当prob极小(接近浮点数精度下限)时,-log(runif(n)) / prob会变成极大值甚至Inf,但实际计算中,由于浮点数精度限制,小权重样本的该值会被大权重样本完全压制,导致永远无法进入前size的排序结果,出现"截断"现象。
此外,权重缩放无法解决该问题,因为缩放只是对prob进行倍数调整,相对比例不变,小权重的-log(runif)/prob依然会被大权重的数值覆盖。
解决方案
1. 使用dplyr::slice_sample()替代
dplyr包的slice_sample()函数在无放回带权重采样时,采用的是逐次抽样并重新加权的逻辑,和手动实现一致,能正确处理小概率样本:
library(dplyr) sample_dplyr = function(N_) { dat %>% slice_sample(n = N_, weight_by = lkhd_wt, replace = FALSE) %>% pull(id) }
2. 使用sampling包的不等概率无放回采样函数
sampling包提供了专门的不等概率无放回采样实现,比如UPbrewer()函数,适合处理此类场景:
library(sampling) sample_sampling = function(N_) { idx = UPbrewer(dat$lkhd_wt, N_) dat$id[idx] }
3. 手动实现优化
如果不想依赖第三方包,可以优化手动采样逻辑,提升效率(比如避免循环中重复计算权重):
sample_manual_opt = function(N_) { remaining_ids = dat$id remaining_wts = dat$lkhd_wt smpl_ = integer(N_) for(i in 1:N_) { norm_wts = remaining_wts / sum(remaining_wts) pick_idx = sample(length(remaining_ids), size = 1, prob = norm_wts) smpl_[i] = remaining_ids[pick_idx] remaining_ids = remaining_ids[-pick_idx] remaining_wts = remaining_wts[-pick_idx] } smpl_ }
内容的提问来源于stack exchange,提问作者Ryan S
相关产品推荐
相关产品推荐

