如何在PySpark中过滤含无重复元素ArrayType列的数据集
PySpark过滤不含重复元素的数组行
方法一:比较原数组与去重后数组的长度
核心逻辑:若数组去重后的长度和原长度一致,说明数组内无重复元素,保留该行。
- 创建示例数据集
from pyspark.sql import SparkSession spark = SparkSession.builder.appName("FilterArrayDuplicates").getOrCreate() data = [ ("a", [1,2]), ("a", [2,2]), ("a", [1,3]), ("a", [1,2,3]), ("a", [1,1,3]) ] df = spark.createDataFrame(data, schema=["A", "B"]) df.show()
- 执行过滤操作
from pyspark.sql.functions import size, distinct filtered_df = df.filter(size(df.B) == size(distinct(df.B))) filtered_df.show()
这种方法简单直观,适合大多数常规场景。
方法二:使用aggregate函数提前终止重复检查
针对元素较多的数组,该方法会在遍历过程中一旦发现重复就停止检查,性能更优:
from pyspark.sql.functions import expr filtered_df = df.filter(expr(""" aggregate( B, (seen = array(), has_dup = false), (acc, x) -> if(acc.has_dup, acc, (array_union(acc.seen, array(x)), array_contains(acc.seen, x))), acc -> not(acc.has_dup) ) """)) filtered_df.show()
逻辑说明:遍历数组时维护一个已见元素集合和重复标记,发现元素已存在则标记为有重复,最终仅保留无重复的行。
内容的提问来源于stack exchange,提问作者Programmer
相关产品推荐
相关产品推荐

