如何优化R data.table多列分组键相交的大数据处理效率?
优化data.table分组内特征筛选的性能问题
重要编辑:因现有回答有误,特此澄清
我手头有个data.table,里面包含分组列(split_by)、键列(key_by)和特征ID列(intersect_by)。我的需求很明确:在每个split_by分组里,只保留那些特征ID被该分组内所有键都共享的行。
举个具体的例子,我构造了这样的数据集:
dt <- data.table(id = 1:6, key1 = 1, key2 = c(1:2, 2), group_id1= 1, group_id2= c(1:2, 2:1, 1:2), trait_id1 = 1, trait_id2 = 2:1) setkey(dt, group_id1, group_id2, trait_id1, trait_id2) dt
输出如下:
id key1 key2 group_id1 group_id2 trait_id1 trait_id2 1: 4 1 1 1 1 1 1 2: 1 1 1 1 1 1 2 3: 5 1 2 1 1 1 2 4: 2 1 2 1 2 1 1 5: 6 1 2 1 2 1 1 6: 3 1 2 1 2 1 2
我期望得到的结果res是这样的:
> res[] id key1 key2 group_id1 group_id2 trait_id1 trait_id2 1: 1 1 1 1 1 1 2 2: 5 1 2 1 1 1 2 3: 2 1 2 1 2 1 1 4: 6 1 2 1 2 1 1 5: 3 1 2 1 2 1 2
这里id为4的行被剔除了,原因是它所在的group_id1=1、group_id2=1分组里有(1,1)和(1,2)两种键组合,而特征(1,1)只被其中一种键包含,没被所有键共享;反过来,id为1和5的行对应的特征被该分组所有键共享,所以保留了下来。
我自己写了一个实现函数:
intersect_this_by2 <- function(dt, key_by = NULL, split_by = NULL, intersect_by = NULL){ dtc <- as.data.table(dt) # 计算分组内的键数量 dtc[, n_keys := uniqueN(.SD), by = split_by, .SDcols = key_by] # 计算每个分组中各特征对应的键数量,仅保留与分组总键数相等的行 dtc[, keep := n_keys == uniqueN(.SD), by = c(intersect_by, split_by), .SDcols = key_by] dtc <- dtc[keep == TRUE][, c("n_keys", "keep") := NULL] return(dtc) }
但这个函数在处理大数据集的时候(比如1000万行、30个特征水平)速度特别慢,想问问有没有优化方案?另外这个函数有没有什么明显的缺陷?
最终编辑: 后来Uwe给出了一个更简洁的优化方案,比原代码快了40%,实现函数如下:
intersect_this_by_uwe <- function(dt, key_by = c("key1"), split_by = c("group_id1", "group_id2"), intersect_by = c("trait_id1", "trait_id2")){ dti <- copy(dt) dti[, original_order_id__ := 1:.N] setkeyv(dti, c(split_by, intersect_by, key_by)) uni <- unique(dti, by = c(split_by, intersect_by, key_by)) unique_keys_by_group <- unique(uni, by = c(split_by, key_by))[, .N, by = c(split_by)] unique_keys_by_group_and_trait <- uni[, .N, by = c(split_by, intersect_by)] # 第一次连接筛选出键数量匹配的分组/特征组合 selected_groups_and_traits <- unique_keys_by_group_and_trait[unique_keys_by_group, on = c(split_by, "N"), nomatch = 0L] # 第二次连接筛选出有效子集的记录 dti[selected_groups_and_traits, on = c(split_by, intersect_by)][ order(original_order_id__), -c("original_order_id__","N")] }
我们在1000万行的数据集上做了基准测试,结果如下:
> microbenchmark::microbenchmark(old_way = {res <- intersect_this_by(dt, + key_by = c("key1"), + split_by = c("group_id1", "group_id2"), + intersect_by = c("trait_id1", "trait_id2"))}, + new_way = {res <- intersect_this_by2(dt, + key_by = c("key1"), + split_by = c("group_id1", "group_id2"), + intersect_by = c("trait_id1", "trait_id2"))}, + new_way_uwe = {res <- intersect_this_by_uwe(dt, + key_by = c("key1"), + split_by = c("group_id1", "group_id2"), + intersect_by = c("trait_id1", "trait_id2"))}, + times = 10) Unit: seconds expr min lq mean median uq max neval cld old_way 3.145468 3.530898 3.514020 3.544661 3.577814 3.623707 10 b new_way 15.670487 15.792249 15.948385 15.988003 16.097436 16.206044 10 c new_way_uwe 1.982503 2.350001 2.320591 2.394206 2.412751 2.436381 10 a
内容的提问来源于stack exchange,提问作者BenoitLondon
相关产品推荐
相关产品推荐

