PySpark中基于flag_code列值变化获取首个后续日期的方法
PySpark实现状态变更首个next_date的简便方法
需求说明
找到flag_code列状态(0或1)发生变化时的首个next_date;当连续出现多个相同状态时,这些行对应同一个首个状态变更日期。
输入示例
| id | flag_code | date |
|---|---|---|
| 1 | 0 | 2022-10-01 |
| 1 | 1 | 2022-10-02 |
| 1 | 0 | 2022-10-03 |
| 1 | 0 | 2022-10-04 |
| 1 | 0 | 2022-11-20 |
| 1 | 1 | 2023-02-01 |
期望输出
| id | flag_code | date | next_date |
|---|---|---|---|
| 1 | 0 | 2022-10-01 | 2022-10-02 |
| 1 | 1 | 2022-10-02 | 2022-10-03 |
| 1 | 0 | 2022-10-03 | 2023-02-01 |
| 1 | 0 | 2022-10-04 | 2023-02-01 |
| 1 | 0 | 2022-11-20 | 2023-02-01 |
| 1 | 1 | 2023-02-01 | NULL |
实现步骤
使用PySpark窗口函数和分组聚合即可实现,具体代码如下:
- 导入依赖
from pyspark.sql import Window import pyspark.sql.functions as F
- 创建基础排序窗口,按
id分区、date升序排列
order_window = Window.partitionBy("id").orderBy("date")
- 标记连续相同状态的分组
通过比较当前行与前一行的flag_code,标记状态变更点,再累加得到连续状态的分组ID:
df = df.withColumn( "change_flag", F.when(F.lag("flag_code").over(order_window) != F.col("flag_code"), 1).otherwise(0) ).withColumn( "group_id", F.sum("change_flag").over(order_window.rowsBetween(Window.unboundedPreceding, 0)) )
- 计算每个分组对应的首个变更日期
先按id和group_id分组获取每组的最大日期,再用lead函数获取下一个分组的最大日期作为当前组的next_date,最后关联回原表:
# 获取每个分组对应的next_date group_next_df = df.groupBy("id", "group_id")\ .agg(F.max("date").alias("group_max_date"))\ .withColumn( "next_date", F.lead("group_max_date").over(Window.partitionBy("id").orderBy("group_max_date")) ) # 关联回原表,填充所有行的next_date result_df = df.join(group_next_df, on=["id", "group_id"], how="left")\ .select("id", "flag_code", "date", "next_date")
执行上述代码后,即可得到符合需求的结果。
内容的提问来源于stack exchange,提问作者Jresearcher
相关产品推荐
相关产品推荐

