Spark DataFrame如何按数量阈值拆分groupBy聚合结果为多行
Spark分组后按阈值拆分行实现方案
需求场景
按attr1字段分组聚合attr2值,单组内聚合得到的值列表每超过3个就拆分为新行,每行最多保留3个逗号拼接的值,适配单次最多传入固定数量参数的API调用要求。
现有基础聚合代码:
val tdf1=tDf.groupBy(col("attr1")).agg(concat_ws(",",collect_list(col("attr2"))).as("value"))
当前聚合输出:
scala> tdf1.show(false) +-----+-----------------------+ |attr1|value | +-----+-----------------------+ |1 |2,20 | |2 |200,201,202,203,204,205| +-----+-----------------------+
期望输出:
scala> tdf1.show(false) +-----+-----------------------+ |attr1|value | +-----+-----------------------+ |1 |2,20 | |2 |200,201,202 | |2 |203,204,205 | +-----+-----------------------+
原始输入数据结构:
[{ "attr1":"1", "attr2":"2" },{ "attr1":"1", "attr2":"20" },{ "attr1":"2", "attr2":"200" },{ "attr1":"2", "attr2":"201" },{ "attr1":"2", "attr2":"202" },{ "attr1":"2", "attr2":"203" },{ "attr1":"2", "attr2":"204" },{ "attr1":"2", "attr2":"205" }]
实现代码
推荐直接对分组收集的数组做切片拆分,避免先拼接字符串再拆分的冗余性能损耗,阈值可根据API实际入参限制灵活调整:
import org.apache.spark.sql.functions._ // 配置单组最大保留值数量,和API入参限制对齐 val maxItemPerRow = 3 val result = tDf .groupBy("attr1") // 先收集组内所有attr2为数组,不提前拼接字符串 .agg(collect_list("attr2").as("attr2_arr")) // 按阈值生成切片数组:每maxItemPerRow个元素切为一个子数组 .withColumn("sliced_arr", expr(s"transform(sequence(0, size(attr2_arr)-1, $maxItemPerRow), idx -> slice(attr2_arr, idx+1, $maxItemPerRow))") ) // 炸开切片数组,每个子数组对应独立一行 .withColumn("item_per_row", explode("sliced_arr")) // 子数组拼接为逗号分隔字符串,和预期输出格式对齐 .withColumn("value", concat_ws(",", col("item_per_row"))) .select("attr1", "value")
执行result.show(false)即可得到完全匹配预期的输出结果。
注意事项
- 如果需要保证
attr2的拼接顺序,可以在分组前先对DataFrame按业务要求排序,或者对收集到的attr2_arr使用sort_array函数排序后再做切片 - 仅需修改
maxItemPerRow参数即可适配不同入参数量限制的API,不需要调整核心逻辑
内容的提问来源于stack exchange,提问作者Saawan
相关产品推荐
相关产品推荐

