PySpark按ID分组聚合连续pos对应的pos与value数组
Spark按ID分组并聚合连续pos的数组
原始DataFrame
+---+-----+------------+ | ID| pos| value| +---+-----+------------+ | A1| 0| ABC| | A1| 1| BCD| | A1| 2| CDF| | A1| 5| ABC| | A1| 8| FGR| | A1| 9| EFD| | A2| 0| L_1| | A2| 1| L_2| | A2| 3| STU| +---+-----+------------+
期望输出
+---+-------------+-----------------------------+ | ID| arr(pos)| arr(value)| +---+-------------+-----------------------------+ | A1| [0, 1, 2]| ['ABC', 'BCD', 'CDF']| | A1| [5]| ['ABC']| | A1| [8, 9]| ['FGR', 'EFD']| | A2| [0, 1]| ['L_1', 'L_2']| | A2| [3]| ['STU']| +---+-------------+-----------------------------+
核心需求
- 按
ID分组,对pos和value分别聚合为数组 - 同一
ID下,pos数组中的数值必须是连续整数;若pos与前后数值不连续(如A1的pos=5),则单独成组 - 同一
ID的pos值无重复
已尝试的方法
方法1:使用Lag函数识别前后pos
代码:
from pyspark.sql import functions as F, Window as W window = W.partitionBy("ID").orderBy("pos") df = ( df .withColumn("prev_pos", F.lag(F.col('pos'), 1).over(window)) .withColumn("next_pos", F.lead(F.col('pos'), 1).over(window)) # 注:原代码写错,lag(-1)应为lead )
执行结果:
+---+-----+------------+---------+---------+ | ID| pos| value| prev_pos| next_pos| +---+-----+------------+---------+---------+ | A1| 0| ABC| null| 1| | A1| 1| BCD| 0| 2| | A1| 2| CDF| 1| 5| | A1| 5| ABC| 2| 8| | A1| 8| FGR| 5| 9| | A1| 9| EFD| 8| null| | A2| 0| L_1| null| 1| | A2| 1| L_2| 0| 3| | A2| 3| STU| 1| null| +---+-----+------------+---------+---------+
问题:无法进一步识别连续区间并分组聚合
方法2:直接分组聚合所有pos
代码:
df = ( df .groupby("ID") .agg( F.sort_array(F.collect_list(F.struct("pos", "value"))) .alias("pos_list") ) .withColumn("arr(pos)", F.col("pos_list.pos")) .withColumn("arr(value)", F.col("pos_list.value")) .drop("pos_list") )
执行结果:
+---+--------------------+---------------------------------------------+ | ID| arr(pos)| arr(value)| +---+--------------------+---------------------------------------------+ | A1| [0, 1, 2, 5, 8, 9]| ['ABC', 'BCD', 'CDF', 'ABC', 'FGR', 'EFD']| | A2| [0, 1, 3]| ['L_1', 'L_2', 'STU']| +---+--------------------+---------------------------------------------+
问题:无法在pos不连续处拆分数组
解决方案
核心思路:先为每个连续的pos区间标记分组ID,再按ID和区间分组ID聚合。
完整代码:
from pyspark.sql import functions as F, Window as W # 1. 定义窗口:按ID分区,pos排序 window = W.partitionBy("ID").orderBy("pos") # 2. 标记连续区间的起点:第一个元素或与前一个pos差值≠1的位置 df_with_gap = df.withColumn( "is_new_group", F.when( F.lag("pos").over(window).isNull(), 1 ).when( F.col("pos") - F.lag("pos").over(window) != 1, 1 ).otherwise(0) ) # 3. 累积求和生成区间分组ID df_with_group = df_with_gap.withColumn( "group_id", F.sum("is_new_group").over(window.rangeBetween(W.unboundedPreceding, 0)) ) # 4. 按ID和group_id聚合,收集数组 result = df_with_group.groupBy("ID", "group_id").agg( F.collect_list("pos").alias("arr(pos)"), F.collect_list("value").alias("arr(value)") ).drop("group_id").orderBy("ID", F.min("pos").over(W.partitionBy("ID"))) # 展示结果 result.show(truncate=False)
执行结果与期望输出一致:
+---+---------+----------------+ |ID |arr(pos) |arr(value) | +---+---------+----------------+ |A1 |[0,1,2] |[ABC, BCD, CDF] | |A1 |[5] |[ABC] | |A1 |[8,9] |[FGR, EFD] | |A2 |[0,1] |[L_1, L_2] | |A2 |[3] |[STU] | +---+---------+----------------+
内容的提问来源于stack exchange,提问作者huy
相关产品推荐
相关产品推荐

