PySpark实现Dataframe按10秒滑动窗口分组(超间隔重置)
基于动态10秒窗口的Spark分组实现
需求说明
需要将以下输入DataFrame按动态10秒窗口分组:当当前记录与同组第一条记录的时间差不超过10秒时归为一组,超过则重置分组计数器;而非采用固定时间切片(如00-10秒、10-20秒这类硬划分区间)。
输入数据
+----+-------------------+ | id | date| +----+-------------------+ | A|2016-03-11 09:00:00| | B|2016-03-11 09:00:07| | C|2016-03-11 09:00:18| | D|2016-03-11 09:00:21| | E|2016-03-11 09:00:39| | F|2016-03-11 09:00:44| | G|2016-03-11 09:00:49| +----+-------------------+
预期输出
+----+-------------------+-----+ | id | date|group| +----+-------------------+-----+ | A|2016-03-11 09:00:00| 1 | | B|2016-03-11 09:00:07| 1 | | C|2016-03-11 09:00:18| 2 | | D|2016-03-11 09:00:21| 2 | | E|2016-03-11 09:00:39| 3 | | F|2016-03-11 09:00:44| 3 | | G|2016-03-11 09:00:49| 4 | +----+-------------------+-----+
固定时间切片的问题
固定时间切片会按预设区间硬划分,导致不符合需求的结果:
+----+-------------------+-----+ | id | date|group| +----+-------------------+-----+ | A|2016-03-11 09:00:00| 1 | | B|2016-03-11 09:00:07| 1 | | C|2016-03-11 09:00:18| 2 | | D|2016-03-11 09:00:21| 3 | | E|2016-03-11 09:00:39| 4 | | F|2016-03-11 09:00:44| 5 | | G|2016-03-11 09:00:49| 5 | +----+-------------------+-----+
解决方案
实现思路
- 确保数据按时间排序(动态分组依赖记录的时间顺序)
- 计算当前记录与前一条记录的时间差
- 标记时间差超过10秒的行作为新分组的触发点
- 对触发点累加生成连续的分组ID
代码实现
from pyspark.sql import SparkSession from pyspark.sql.window import Window from pyspark.sql.functions import col, lag, unix_timestamp, sum as spark_sum # 初始化SparkSession spark = SparkSession.builder.appName("DynamicTimeGrouping").getOrCreate() # 构造输入DataFrame data = [ ("A", "2016-03-11 09:00:00"), ("B", "2016-03-11 09:00:07"), ("C", "2016-03-11 09:00:18"), ("D", "2016-03-11 09:00:21"), ("E", "2016-03-11 09:00:39"), ("F", "2016-03-11 09:00:44"), ("G", "2016-03-11 09:00:49") ] df = spark.createDataFrame(data, ["id", "date"]) # 将字符串类型的date转为timestamp类型 df = df.withColumn("date", col("date").cast("timestamp")) # 定义窗口:按时间排序,用于获取前一行记录 time_window = Window.orderBy("date") # 计算当前行与前一行的时间差(秒),第一行无前置记录,时间差设为0 df = df.withColumn("prev_date", lag("date", 1).over(time_window)) df = df.withColumn("time_diff", unix_timestamp("date") - unix_timestamp("prev_date")) df = df.fillna({"time_diff": 0}) # 标记分组触发点:时间差超过10秒则标记为1,否则为0 df = df.withColumn("group_flag", (col("time_diff") > 10).cast("int")) # 累加触发点得到分组ID(加1是为了让分组从1开始计数) df = df.withColumn("group", spark_sum("group_flag").over(time_window) + 1) # 选择最终需要的列并展示结果 result_df = df.select("id", "date", "group") result_df.show()
逻辑说明
lag函数用于获取前一条记录的时间,结合unix_timestamp计算时间差- 时间差超过10秒的行标记为1,代表新分组的起始
- 对标记点进行累加,得到连续的分组ID,加1是为了避免分组从0开始
- 整个过程基于记录的时间顺序,确保分组是动态的,而非固定区间划分
内容的提问来源于stack exchange,提问作者Voxeldoodle
相关产品推荐
相关产品推荐

