PySpark按行值过滤/拆分含STOP标识的混合数据
解决方案:PySpark/Spark SQL 处理分段数据过滤
核心思路
要解决这个问题,关键是先给两段STOP(示例中的0)之间的片段分配唯一分组ID,再统计每个分组内有效元素(非STOP)的数量,最后保留有效元素数量达标(≥min_length)的分组所有行。
因为Spark分布式特性无法用本地计数器,所以用窗口函数实现分组ID的累计计算,再通过分组统计和关联筛选完成过滤。
PySpark 实现
from pyspark.sql import functions as F from pyspark.sql.window import Window # 1. 构建示例DataFrame df = spark.createDataFrame([(1,),(0,),(3,),(0,),(4,),(4,),(5,),(0,)], schema=["value"]) # 2. 定义窗口(注意:生产环境需替换为真实排序键,比如业务自增ID/时间戳) window_spec = Window.orderBy(F.monotonically_increasing_id()) # 3. 为每行分配分组ID:统计当前行之前出现的STOP(0)次数作为组ID df_with_group = df.withColumn( "group_id", F.sum(F.when(F.col("value") == 0, 1).otherwise(0)).over(window_spec.rowsBetween(Window.unboundedPreceding, Window.currentRow - 1)) ) # 4. 统计每个分组的有效元素数量 group_stats = df_with_group.groupBy("group_id").agg( F.count(F.when(F.col("value") != 0, 1)).alias("non_zero_count") ) # 5. 筛选并保留达标分组的所有行 min_length = 2 result_df = df_with_group.join( group_stats.filter(F.col("non_zero_count") >= min_length), on="group_id", how="inner" ).drop("group_id", "non_zero_count") # 查看结果 result_df.show()
输出结果:
+-----+ |value| +-----+ | 4| | 4| | 5| | 0| +-----+
Spark SQL 实现
-- 1. 注册临时视图 CREATE OR REPLACE TEMP VIEW data AS SELECT value FROM VALUES (1),(0),(3),(0),(4),(4),(5),(0) AS t(value); -- 2. 设置最小保留长度 SET min_length = 2; -- 3. 执行分段过滤逻辑 WITH data_with_group AS ( SELECT value, SUM(CASE WHEN value = 0 THEN 1 ELSE 0 END) OVER ( ORDER BY monotonically_increasing_id() ROWS BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW - 1 ) AS group_id FROM data ), group_stats AS ( SELECT group_id, COUNT(CASE WHEN value != 0 THEN 1 END) AS non_zero_count FROM data_with_group GROUP BY group_id ) SELECT d.value FROM data_with_group d JOIN group_stats gs ON d.group_id = gs.group_id WHERE gs.non_zero_count >= ${min_length} ORDER BY monotonically_increasing_id();
Spark-Shell(Scala)实现
import org.apache.spark.sql.functions._ import org.apache.spark.sql.expressions.Window // 1. 构建示例DataFrame val df = spark.createDataFrame(Seq((1,),(0,),(3,),(0,),(4,),(4,),(5,),(0,))).toDF("value") // 2. 定义窗口(生产环境替换为真实排序键) val windowSpec = Window.orderBy(monotonically_increasing_id()) // 3. 分配分组ID val dfWithGroup = df.withColumn( "group_id", sum(when(col("value") === 0, 1).otherwise(0)).over(windowSpec.rowsBetween(Window.unboundedPreceding, Window.currentRow - 1)) ) // 4. 统计分组有效元素数量并筛选达标组 val minLength = 2 val groupStats = dfWithGroup.groupBy("group_id").agg( count(when(col("value") =!= 0, 1)).alias("non_zero_count") ).filter(col("non_zero_count") >= minLength) // 5. 关联得到结果 val resultDf = dfWithGroup.join(groupStats, Seq("group_id"), "inner").drop("group_id", "non_zero_count") // 查看结果 resultDf.show()
注意事项
- 示例中用
monotonically_increasing_id()作为排序键仅用于演示,生产环境必须使用业务侧的有序字段(如时间戳、自增ID),否则分布式环境下数据顺序会混乱,导致分组错误。 - 若最后一段数据没有STOP标识,可调整分组统计逻辑,判断是否为最后一个分组并按需保留。
内容的提问来源于stack exchange,提问作者Syrius
相关产品推荐
相关产品推荐

