如何高效使用R data.table按超集聚合更新每行数据?
问题描述
给定如下data.table数据表:
dt = data.table(id = 1:5, intMask = c(11,14,8,1,13), imprint = c("1011", "1110", "1000", "0001", "1101"), N = c(3,3,1,1,3), mass = c(.05,.1,.15,.3,.4))
对应的数据表内容:
id intMask imprint N mass 1: 1 11 1011 3 0.05 2: 2 14 1110 3 0.10 3: 3 8 1000 1 0.15 4: 4 1 0001 1 0.30 5: 5 13 1101 3 0.40
字段说明:
imprint:基数为5的集合子集的二进制表示intMask:对应二进制的整数形式N:集合基数(即二进制字符串中1的个数)
需求:计算newMass列,值为当前行所有超集行的mass之和。超集判定规则:若行A的intMask与行B的intMask按位与结果等于行B的intMask,且行A的N≥行B的N,则行A是行B的超集。
当前使用循环实现的代码如下,但数据量达到数千行时效率极低:
for(i in 1:nrow(dt)) { i.intMask <- dt[i,intMask] i.N <- dt[i,N] dt[i, newMass := sum(dt[N >= i.N,][bitwAnd(intMask, i.intMask) == i.intMask, mass])] }
期望输出(注:原示例中第5行newMass的0.04为笔误,正确值应为0.40):
id intMask imprint N mass newMass 1: 1 11 1011 3 0.05 0.05 2: 2 14 1110 3 0.10 0.10 3: 3 8 1000 1 0.15 0.70 4: 4 1 0001 1 0.30 0.75 5: 5 13 1101 3 0.40 0.40
高效实现方案(基于data.table)
核心思路
循环方案的低效根源在于逐行重复筛选和计算,我们可以利用data.table的非等连接+按位运算特性,一次性完成所有行的newMass计算:
- 用自身连接缩小匹配范围:以
N >= N为连接条件,先筛选出所有可能的超集候选行(超集的集合基数必然≥子集) - 按位运算过滤真正的超集:对候选行执行
bitwAnd(intMask, i.intMask) == i.intMask,保留符合超集规则的行 - 分组聚合求和:按原表的
id分组,对匹配到的mass求和得到newMass - 将结果更新回原表
实现代码
# 执行自身连接、过滤超集并聚合计算 agg_result <- dt[dt, on = .(N >= N), allow.cartesian = TRUE ][bitwAnd(intMask, i.intMask) == i.intMask, .(newMass = sum(mass)), by = .(i.id)] # 将聚合结果更新回原数据表 dt[agg_result, newMass := i.newMass, on = .(id = i.id)]
结果验证
运行上述代码后,dt的输出与期望结果一致:
id intMask imprint N mass newMass 1: 1 11 1011 3 0.05 0.05 2: 2 14 1110 3 0.10 0.10 3: 3 8 1000 1 0.15 0.70 4: 4 1 0001 1 0.30 0.75 5: 5 13 1101 3 0.40 0.40
性能优势
- 利用
data.table的C级连接引擎替代R循环,大幅降低计算开销 - 先通过
N的非等连接缩小候选范围,再执行按位运算过滤,减少无效计算 - 聚合操作批量完成,避免了循环中重复的数据子集提取
内容的提问来源于stack exchange,提问作者Vaclav Kratochvíl
相关产品推荐
相关产品推荐

