sparklyr中如何计算两个列表类型列的交集及交集大小
sparklyr计算两个数组列的交集及交集长度
需求说明
操作tbl_spark对象时,针对包含两个数组(列表)类型列的Spark DataFrame,需要输出两项结果:
- 两个数组的交集,返回结果为数组(列表)类型
- 交集中包含的元素总数
测试数据构造代码如下:
library(dplyr) library(sparklyr) ## 已提前创建Spark连接对象sc mtcars_spark <- copy_to(sc, mtcars) ## 构造含列表列的测试表 tbl_with_lists <- mtcars_spark %>% mutate(mpg_rounded = round(mpg, -1)) %>% group_by(mpg_rounded) %>% summarize( cyl_all = paste(collect_set(as.character(cyl)), sep = ", "), gear_all = paste(collect_set(as.character(gear)), sep = ", ") ) %>% ungroup() %>% ft_regex_tokenizer("cyl_all", "cyl_list", pattern = "[,]\\s*") %>% ft_regex_tokenizer("gear_all", "gear_list", pattern = "[,]\\s*")
测试表预览:
# Source: spark<?> [?? x 5] mpg_rounded cyl_all gear_all cyl_list gear_list <dbl> <chr> <chr> <list> <list> 1 10 8.0 3.0 <list [1]> <list [1]> 2 30 4.0 5.0, 4.0 <list [1]> <list [2]> 3 20 8.0, 6.0, 4.0 5.0, 3.0, 4.0 <list [3]> <list [3]>
实现方案
优先使用Spark内置原生函数实现,全程分布式执行,性能最优,不需要编写自定义UDF,也不需要把数据拉取到本地计算。
Spark 2.4及以上版本内置array_intersect函数,直接传入两个数组列即可返回交集,搭配size函数可直接统计交集元素个数,代码如下:
result <- tbl_with_lists %>% mutate( # 计算两个列表的交集,返回列表类型 intersect_list = sql("array_intersect(cyl_list, gear_list)"), # 统计交集元素总数 intersect_count = sql("size(intersect_list)") )
结果验证
针对上述测试数据,运行后返回结果符合预期:
- mpg_rounded=10:cyl_list为
["8.0"],gear_list为["3.0"],交集为空列表,交集计数为0 - mpg_rounded=30:cyl_list为
["4.0"],gear_list为["5.0", "4.0"],交集为["4.0"],交集计数为1 - mpg_rounded=20:cyl_list为
["8.0", "6.0", "4.0"],gear_list为["5.0", "3.0", "4.0"],交集为["4.0"],交集计数为1
低版本兼容方案
如果使用的Spark版本低于2.4,可直接调用Hive内置的array_intersect函数,写法与上述代码完全一致,仅需确保当前Spark会话开启了Hive支持(sparklyr默认创建的本地连接默认开启)。不推荐使用数组explode后join聚合的实现方式,该方法shuffle开销大,性能远低于内置函数方案。
注意:通过
ft_regex_tokenizer生成的列表列本身就是Spark的ArrayType类型,完全匹配array_intersect的入参要求,不需要额外做类型转换。
内容的提问来源于stack exchange,提问作者MRipley
相关产品推荐
相关产品推荐

