PySpark:统计特定station_no的连续1值streak编号
解决方案:Spark中按分组统计连续1值的递增序列
这个问题核心是在每个station_no分组内,识别连续的group_flag=1区块,为每个区块分配递增编号,group_flag=0的行保持原值。可以通过Spark窗口函数结合累加操作实现,具体步骤如下:
实现思路
- 标记连续1的起始行:按
station_no分区、period_number排序,用lag函数获取前一行的group_flag,判断当前行是否为连续1区块的起点(当前group_flag=1且前一行是0,或是分组第一行且group_flag=1)。 - 累加起始点生成区块编号:对分组内的起始点标记进行累加,得到每个连续1区块的唯一递增编号。
- 替换原
group_flag值:保留group_flag=0的行,其余行用累加得到的区块编号替换。
完整代码实现
from pyspark.sql import Window from pyspark.sql.functions import col, lag, sum, when # 定义基础窗口:按station_no分区,period_number升序排序 base_window = Window.partitionBy("station_no").orderBy("period_number") # 步骤1:标记连续1区块的起始行 df_with_start = df.withColumn( "is_start", when( (col("group_flag") == 1) & (lag("group_flag", 1, 0).over(base_window) == 0), 1 ).otherwise(0) ) # 步骤2:累加起始点,生成每个连续区块的递增编号 cum_window = Window.partitionBy("station_no").orderBy("period_number").rowsBetween(Window.unboundedPreceding, Window.currentRow) df_with_streak = df_with_start.withColumn( "streak_id", sum("is_start").over(cum_window) ) # 步骤3:替换原group_flag,保留0值,1值替换为区块编号 result_df = df_with_streak.withColumn( "group_flag", when(col("group_flag") == 0, 0).otherwise(col("streak_id")) ).drop("is_start", "streak_id") # 查看最终结果 result_df.show()
运行结果
执行代码后,输出将与你期望的结果一致:
+----------+-------------+----------+ |station_no|period_number|group_flag| +----------+-------------+----------+ | BAN| 0| 1| | BAN| 1| 0| | BAN| 2| 2| | BAN| 3| 2| | BAN| 4| 2| | BAN| 5| 0| | BAN| 6| 3| | BOZ| 0| 0| | BOZ| 1| 1| | BOZ| 2| 1| | BOZ| 3| 0| | BOZ| 4| 2| | BOZ| 5| 0| | BOZ| 6| 3| +----------+-------------+----------+
代码细节说明
lag函数:获取前一行的group_flag,默认值设为0,处理分组内第一行的边界情况。is_start列:标记当前行是否为新连续1区块的起点,是则为1,否则为0。- 累加操作:在分区内从起始行到当前行累加
is_start,得到的数值就是每个连续区块的递增编号。 when替换:保留原0值行,将1值行替换为对应区块编号,最后删除辅助列。
内容的提问来源于stack exchange,提问作者Bart
相关产品推荐
相关产品推荐

