如何用data.table高效实现按行离散概率抽样?
问题描述
我有一个存储离散分布概率列的data.table,示例代码如下:
dt <- data.table(p1 = c(0.5, 0.25, 0.1), p2 = c(0.25, 0.5, 0.1), p3 = c(0.25, 0.25, 0.8))
希望新增一列,用每行对应的p1/p2/p3概率值抽样得到1-3的随机变量。预想的data.table语法是:
dt[, sample := sample(1:3, 1, prob = c(p1, p2, p3))]
但sample不支持类似pmin/pmax的逐行操作。目前用apply实现了需求,但真实数据集上运行速度很慢,apply实现代码如下:
dt[, sample := apply(dt, 1, function(x) sample(1:3, 1, prob = x[c('p1', 'p2', 'p3')]))]
求data.table的高效实现方式。
高效实现方案
可以通过矢量化累积概率匹配的方式实现,避免逐行循环,充分利用data.table和R内部高效函数的性能:
library(data.table) # 示例数据 dt <- data.table(p1 = c(0.5, 0.25, 0.1), p2 = c(0.25, 0.5, 0.1), p3 = c(0.25, 0.25, 0.8)) # 将概率列转为矩阵,计算每行的累积概率 probs_mat <- as.matrix(dt[, .(p1, p2, p3)]) cum_probs <- t(apply(probs_mat, 1, cumsum)) # 生成每行的随机数,匹配对应的类别 rands <- runif(nrow(dt)) dt[, sample := max.col(rands < cum_probs, ties.method = "first")]
原理说明
- 先将概率列转为矩阵,计算每行的累积概率(比如第一行累积概率为
0.5, 0.75, 1); - 生成与行数相同的均匀随机数
runif(nrow(dt)); - 使用
max.col快速找到每个随机数首次大于累积概率的位置,这个位置就是抽样得到的类别值。
这种方法完全基于矢量化操作,没有逐行循环,在大数据集上的运行速度会比apply实现快一个数量级以上。
如果不想生成中间矩阵,也可以直接在data.table内完成更简洁的写法:
dt[, sample := max.col( runif(.N) < cbind(p1, p1+p2, 1), ties.method = "first" )]
内容的提问来源于stack exchange,提问作者Gregory Apartment
相关产品推荐
相关产品推荐

