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

Java Spark SQL如何按1&&(2||3)规则过滤合并3个同结构数据集?

使用Spark API实现复合规则的数据过滤

过滤需求

需过滤掉满足 规则1 && (规则2 || 规则3) 的数据行,最终保留剩余数据。三个规则定义如下:

  • 规则1:若某col1分组下所有行的pid均为"aaa",则过滤该分组所有行;
  • 规则2:若pid为"aaa"且col1为空,则过滤该行;
  • 规则3:若某(col2,col3)分组下所有行的pid均为"aaa",则过滤该分组所有行;

示例输入数据

pid      col1            col2                 col3  
abaa     111             apple                red
aaa      111             apple                red 
aaa      222             banana               yellow
aaa      222             apple                green  
aaa      333             apple                green
aaa      333             null                 green
aaa      444             apple                red

过滤逻辑说明

  • 等价逻辑:(规则1 && 规则2) || (规则1 && 规则3)
  • 示例中,行aaa 333 null green满足规则1&2,需被过滤;部分行满足规则1&3,需被过滤;

示例输出数据

(注:原输出列名存在笔误,修正为与输入一致的列名)

pid      col1            col2                 col3  
abaa     111             apple                red
aaa      111             apple                red

Spark API 实现方案

Scala 版本代码

import org.apache.spark.sql.functions._

// 模拟示例数据
val data = Seq(
  ("abaa", "111", "apple", "red"),
  ("aaa", "111", "apple", "red"),
  ("aaa", "222", "banana", "yellow"),
  ("aaa", "222", "apple", "green"),
  ("aaa", "333", "apple", "green"),
  ("aaa", "333", null, "green"),
  ("aaa", "444", "apple", "red")
).toDF("pid", "col1", "col2", "col3")

// 计算每个col1分组是否所有pid都是aaa
val col1Group = data.groupBy("col1")
  .agg(count(when(col("pid") =!= "aaa", 1)).alias("non_aaa_count"))
  .withColumn("col1_all_aaa", col("non_aaa_count") === 0)
  .drop("non_aaa_count")

// 计算每个(col2,col3)分组是否所有pid都是aaa
val col2Col3Group = data.groupBy("col2", "col3")
  .agg(count(when(col("pid") =!= "aaa", 1)).alias("non_aaa_count"))
  .withColumn("col2_col3_all_aaa", col("non_aaa_count") === 0)
  .drop("non_aaa_count")

// 关联分组标记到原数据
val joinedData = data.join(col1Group, Seq("col1"), "left")
  .join(col2Col3Group, Seq("col2", "col3"), "left")

// 过滤掉满足规则1&&(规则2||规则3)的行
val filteredData = joinedData.filter(
  !(col("col1_all_aaa") && (
    (col("pid") === "aaa" && col("col1").isNull) || col("col2_col3_all_aaa")
  ))
)

// 查看结果
filteredData.select("pid", "col1", "col2", "col3").show()

Python 版本代码

from pyspark.sql import SparkSession
from pyspark.sql.functions import col, count, when

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

// 模拟示例数据
data = [
    ("abaa", "111", "apple", "red"),
    ("aaa", "111", "apple", "red"),
    ("aaa", "222", "banana", "yellow"),
    ("aaa", "222", "apple", "green"),
    ("aaa", "333", "apple", "green"),
    ("aaa", "333", None, "green"),
    ("aaa", "444", "apple", "red")
]
df = spark.createDataFrame(data, ["pid", "col1", "col2", "col3"])

// 计算col1分组是否全为aaa
col1_group = df.groupBy("col1") \
    .agg(count(when(col("pid") != "aaa", 1)).alias("non_aaa_count")) \
    .withColumn("col1_all_aaa", col("non_aaa_count") == 0) \
    .drop("non_aaa_count")

// 计算(col2,col3)分组是否全为aaa
col2_col3_group = df.groupBy("col2", "col3") \
    .agg(count(when(col("pid") != "aaa", 1)).alias("non_aaa_count")) \
    .withColumn("col2_col3_all_aaa", col("non_aaa_count") == 0) \
    .drop("non_aaa_count")

// 关联分组标记到原数据
joined_df = df.join(col1_group, on="col1", how="left") \
    .join(col2_col3_group, on=["col2", "col3"], how="left")

// 过滤掉满足规则1&&(规则2||规则3)的行
filtered_df = joined_df.filter(
    ~(col("col1_all_aaa") & (
        (col("pid") == "aaa") & col("col1").isNull() | col("col2_col3_all_aaa")
    ))
)

// 展示结果
filtered_df.select("pid", "col1", "col2", "col3").show()

代码逻辑说明

  1. 分组标记计算:通过分组聚合统计每个分组中非"aaa"的pid数量,数量为0则标记该分组所有pid均为"aaa";
  2. 关联标记:将分组计算的标记字段关联回原数据集,使每一行都能获取所属分组的规则匹配状态;
  3. 复合过滤:使用逻辑表达式过滤掉满足规则1 && (规则2 || 规则3)的行,得到最终结果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 19:40:19