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

PySpark:统计特定station_no的连续1值streak编号

解决方案:Spark中按分组统计连续1值的递增序列

这个问题核心是在每个station_no分组内,识别连续的group_flag=1区块,为每个区块分配递增编号,group_flag=0的行保持原值。可以通过Spark窗口函数结合累加操作实现,具体步骤如下:

实现思路

  1. 标记连续1的起始行:按station_no分区、period_number排序,用lag函数获取前一行的group_flag,判断当前行是否为连续1区块的起点(当前group_flag=1且前一行是0,或是分组第一行且group_flag=1)。
  2. 累加起始点生成区块编号:对分组内的起始点标记进行累加,得到每个连续1区块的唯一递增编号。
  3. 替换原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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 14:02:26