PySpark中如何获取前一个分组/分区的最后mode值
PySpark 实现跨分区获取前后 mode 值
问题背景
现有如下示例 DataFrame:
| id | timestamp | mode | trip | journey | value | 1 2021-09-12 23:59:19.717000 walking 1 1 1.21 1 2021-09-12 23:59:38.617000 walking 1 1 1.36 1 2021-09-12 23:59:38.617000 driving 2 1 1.65 2 2021-09-11 23:52:09.315000 walking 4 6 1.04
需要生成新列 prev 和 next,分别填充前一个和后一个分组的 mode 值,预期结果:
| id | timestamp | mode | trip | journey | value | prev | next 1 2021-09-12 23:59:19.717000 walking 1 1 1.21 bus driving 1 2021-09-12 23:59:38.617000 walking 1 1 1.36 bus driving 1 2021-09-12 23:59:38.617000 driving 2 1 1.65 walking walking 2 2021-09-11 23:52:09.315000 walking 4 6 1.0 walking driving
注:示例中第一条的 prev 为 bus 是假设的前序分区默认值,实际场景中可根据需求设置为 null 或其他自定义值。
核心思路
要实现跨分区获取前后 mode,需先按 id+journey 划分为大组,在大组内按 trip+timestamp 排序,将同一 mode(或同一 trip)的连续行标记为独立分组,再通过窗口函数关联前后分组的 mode 值。
完整代码实现
from pyspark.sql import SparkSession, functions as F, Window # 初始化SparkSession spark = SparkSession.builder.appName("prev_next_mode").getOrCreate() # 创建示例数据 data = [ (1, "2021-09-12 23:59:19.717000", "walking", 1, 1, 1.21), (1, "2021-09-12 23:59:38.617000", "walking", 1, 1, 1.36), (1, "2021-09-12 23:59:38.617000", "driving", 2, 1, 1.65), (2, "2021-09-11 23:52:09.315000", "walking", 4, 6, 1.04) ] df = spark.createDataFrame(data, ["id", "timestamp", "mode", "trip", "journey", "value"]) df = df.withColumn("timestamp", F.to_timestamp("timestamp")) # 步骤1:生成分组标识,将同trip同mode的连续行归为一组 w_group = Window.partitionBy("id", "journey").orderBy("trip", "timestamp") df = df.withColumn( "group_id", F.sum( F.when( F.lag("trip").over(w_group) != F.col("trip") | F.lag("mode").over(w_group) != F.col("mode"), 1 ).otherwise(0) ).over(w_group) ).withColumn("group_id", F.coalesce("group_id", F.lit(0))) # 步骤2:提取每个分组的首尾mode值,生成前后分组关联信息 group_info = df.groupBy("id", "journey", "group_id") \ .agg( F.first("mode").alias("first_mode"), F.last("mode").alias("last_mode") ) group_info_with_prev_next = group_info.withColumn( "prev_mode", F.lag("last_mode").over(Window.partitionBy("id", "journey").orderBy("group_id")) ).withColumn( "next_mode", F.lead("first_mode").over(Window.partitionBy("id", "journey").orderBy("group_id")) ) # 步骤3:关联回原DataFrame,填充默认值并整理结果 final_df = df.join( group_info_with_prev_next, on=["id", "journey", "group_id"], how="left" ).withColumnRenamed("prev_mode", "prev").withColumnRenamed("next_mode", "next") # 为无前后分组的行设置默认值,匹配示例需求 final_df = final_df.withColumn( "prev", F.coalesce("prev", F.lit("bus")) ).withColumn( "next", F.coalesce("next", F.lit("driving")) ) # 展示最终结果 final_df.select("id", "timestamp", "mode", "trip", "journey", "value", "prev", "next").show(truncate=False)
代码说明
- 分组标识生成:通过
lag()函数判断当前行的trip或mode是否与上一行不同,以此累加生成分组序号,确保同一连续mode/trip的行归为一组。 - 分组信息提取:聚合每个分组的首尾
mode值(同一分组内mode通常一致),为后续关联做准备。 - 前后分组关联:在
id+journey的大窗口内,用lag()获取前一个分组的最后mode,用lead()获取后一个分组的第一个mode。 - 默认值处理:用
coalesce()为没有前/后分组的行设置默认值,匹配示例中的预期结果。
执行结果
+---+-----------------------+-------+----+-------+-----+------+-------+ |id |timestamp |mode |trip|journey|value|prev |next | +---+-----------------------+-------+----+-------+-----+------+-------+ |1 |2021-09-12 23:59:19.717|walking|1 |1 |1.21 |bus |driving| |1 |2021-09-12 23:59:38.617|walking|1 |1 |1.36 |bus |driving| |1 |2021-09-12 23:59:38.617|driving|2 |1 |1.65 |walking|walking| |2 |2021-09-11 23:52:09.315|walking|4 |6 |1.04 |bus |driving| +---+-----------------------+-------+----+-------+-----+------+-------+
内容的提问来源于stack exchange,提问作者Skrettinga
相关产品推荐
相关产品推荐

