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

如何在PySpark DataFrame中筛选连续3行及以上flag=1的记录并修正值

PySpark 处理连续多行flag=1的序列并重置非目标flag

问题描述

给定按date字段排序的PySpark DataFrame,包含flag列(取值为1或0),需要找出连续3行及以上flag值为1的序列,将不属于这类序列的flag值重置为0。

解决方案

我们可以通过「分组连续相同flag的块」+「统计块长度」的方式实现,比多次嵌套lag/lead更简洁高效,步骤如下:

1. 构造示例数据

from pyspark.sql import SparkSession
from pyspark.sql import functions as F
from pyspark.sql.window import Window

spark = SparkSession.builder.appName("continuous_flag_process").getOrCreate()

# 原始示例数据
raw_data = [
    ("2023-01-01", 1),
    ("2023-01-02", 1),
    ("2023-01-03", 0),
    ("2023-01-04", 1),
    ("2023-01-05", 1),
    ("2023-01-06", 1),
    ("2023-01-07", 1),
    ("2023-01-08", 0),
    ("2023-01-09", 1),
    ("2023-01-10", 0)
]

df = spark.createDataFrame(raw_data, ["date", "flag"]).orderBy("date")

2. 生成连续相同flag的分组ID

用窗口函数标记连续相同flag的块:当当前行flag与前一行不同时,分组ID累加1,以此区分不同的连续序列。

# 按date排序的窗口
sort_window = Window.orderBy("date")

df_with_group = df.withColumn(
    "group_id",
    # 当当前行flag和前一行不同时,加1,否则加0,累加得到分组ID
    F.sum(F.when(F.lag("flag").over(sort_window) != F.col("flag"), 1).otherwise(0)).over(sort_window)
)
# 第一行无前置数据,group_id设为0
df_with_group = df_with_group.withColumn("group_id", F.coalesce(F.col("group_id"), F.lit(0)))

3. 统计每个分组的行数

按分组ID聚合,计算每个连续序列的长度:

group_size_df = df_with_group.groupBy("group_id").agg(F.count("*").alias("group_size"))
df_with_size = df_with_group.join(group_size_df, on="group_id", how="left")

4. 重置flag值

仅保留「flag=1且分组长度≥3」的行的flag为1,其余全部设为0:

result_df = df_with_size.withColumn(
    "new_flag",
    F.when((F.col("flag") == 1) & (F.col("group_size") >= 3), 1).otherwise(0)
).select("date", "flag", "new_flag")

result_df.show()

输出结果

+----------+----+--------+
|      date|flag|new_flag|
+----------+----+--------+
|2023-01-01|   1|       0|
|2023-01-02|   1|       0|
|2023-01-03|   0|       0|
|2023-01-04|   1|       1|
|2023-01-05|   1|       1|
|2023-01-06|   1|       1|
|2023-01-07|   1|       1|
|2023-01-08|   0|       0|
|2023-01-09|   1|       0|
|2023-01-10|   0|       0|
+----------+----+--------+

关于lag函数的执行效率

  • 方案中仅使用了一次lag函数生成分组ID,PySpark的窗口函数是分布式执行的,只要date列有合理的分区策略(如果数据量极大,可以先按业务分区键partitionBy,再在分区内排序),不会有明显性能瓶颈。
  • 相比多次嵌套lag/lead(比如检查当前行+前后两行是否为1),分组统计的方式逻辑更清晰,且仅需一次窗口聚合+一次分组聚合,避免了多次窗口计算的开销,在大数据量场景下性能更优。

内容的提问来源于stack exchange,提问作者joselu_22

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.16 06:15:33