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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 14:30:35