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()
代码逻辑说明
- 分组标记计算:通过分组聚合统计每个分组中非"aaa"的
pid数量,数量为0则标记该分组所有pid均为"aaa"; - 关联标记:将分组计算的标记字段关联回原数据集,使每一行都能获取所属分组的规则匹配状态;
- 复合过滤:使用逻辑表达式过滤掉满足
规则1 && (规则2 || 规则3)的行,得到最终结果。
内容的提问来源于stack exchange,提问作者Sun Yuting
相关产品推荐
相关产品推荐

