PySpark实现按连续时间槽(含重复值)分组客户数据
按连续时间槽分组Spark数据集中的客户
需求说明
需要将Spark数据集中的客户按**连续时间槽(包含重复时间槽)**分组,把连续时间段内的客户归为一组,同时汇总对应的时间槽列表和客户列表。
原始数据集
from pyspark.sql import SparkSession from pyspark.sql import functions as F from pyspark.sql.window import Window spark = SparkSession.builder.appName("continuous_time_slot_grouping").getOrCreate() df = spark.createDataFrame( [(0, 'A'), (1, 'B'), (1, 'C'), (5, 'D'), (8, 'A'), (9, 'F'), (20, 'T'), (20, 'S'), (21, 'C')], ['time_slot', 'customer'])
数据集展示:
+--------+--------+ |time_slot|customer| +--------+--------+ | 0| A| | 1| B| | 1| C| | 5| D| | 8| A| | 9| F| | 20| T| | 20| S| | 21| C| +--------+--------+
期望结果
+--------------------+---------------------------------------------+ | grouped_slots| grouped_customers| +--------------------+---------------------------------------------+ | [0, 1]| ['A', 'B', 'C']| | [5]| ['D']| | [8, 9]| ['A', 'F']| | [20, 21]| ['T', 'S', 'C']| +--------------------+---------------------------------------------+
已尝试步骤
已通过lag函数获取前一条记录的时间槽:
window = Window.orderBy("time_slot") df = df.withColumn("prev_time_slot", F.lag(F.col('time_slot'), 1).over(window))
处理后的数据:
+---------+--------+--------------+ |time_slot|customer|prev_time_slot| +---------+--------+--------------+ | 0| A| null| | 1| B| 0| | 1| C| 1| | 5| D| 1| | 8| A| 5| | 9| F| 8| | 20| T| 9| | 20| S| 20| | 21| C| 20| +---------+--------+--------------+
解决方案
核心逻辑是标记时间槽的不连续边界,生成分组ID后再聚合。
步骤1:标记新分组边界
添加is_new_group列,当当前时间槽与前一个时间槽差值大于1时,标记为新分组的起点:
df = df.withColumn( "is_new_group", F.when(F.col("prev_time_slot").isNull(), 1) .when(F.col("time_slot") - F.col("prev_time_slot") > 1, 1) .otherwise(0) )
处理后的数据:
+---------+--------+--------------+-----------+ |time_slot|customer|prev_time_slot|is_new_group| +---------+--------+--------------+-----------+ | 0| A| null| 1| | 1| B| 0| 0| | 1| C| 1| 0| | 5| D| 1| 1| | 8| A| 5| 1| | 9| F| 8| 0| | 20| T| 9| 1| | 20| S| 20| 0| | 21| C| 20| 0| +---------+--------+--------------+-----------+
步骤2:生成分组ID
对is_new_group列做累加求和,得到每行对应的分组ID:
df = df.withColumn( "group_id", F.sum("is_new_group").over(window.rangeBetween(Window.unboundedPreceding, 0)) )
处理后的数据:
+---------+--------+--------------+-----------+--------+ |time_slot|customer|prev_time_slot|is_new_group|group_id| +---------+--------+--------------+-----------+--------+ | 0| A| null| 1| 1| | 1| B| 0| 0| 1| | 1| C| 1| 0| 1| | 5| D| 1| 1| 2| | 8| A| 5| 1| 3| | 9| F| 8| 0| 3| | 20| T| 9| 1| 4| | 20| S| 20| 0| 4| | 21| C| 20| 0| 4| +---------+--------+--------------+-----------+--------+
步骤3:按分组ID聚合
最后按group_id分组,聚合得到时间槽列表和客户列表:
result_df = df.groupBy("group_id").agg( F.collect_set("time_slot").alias("grouped_slots"), F.collect_list("customer").alias("grouped_customers") ).drop("group_id") # 可选:对时间槽列表排序,保证顺序正确 result_df = result_df.withColumn("grouped_slots", F.sort_array(F.col("grouped_slots"))) result_df.show(truncate=False)
运行后即可得到期望的输出。
内容的提问来源于stack exchange,提问作者Islam Elbanna
相关产品推荐
相关产品推荐

