You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.21 18:44:53