PySpark中按条件保留重复数据的实现方法
解决Spark DataFrame分组过滤问题
需求说明
给定Spark DataFrame,需按rank列分组处理:
- 若分组内
group列所有值完全相同且分组行数≥2,移除该分组所有行 - 若分组内
group列存在多种值,或分组仅一行(无论group值),保留该分组所有行
原始数据
data = [['p1', 0, 'dog'], ['p2', 0, 'dog'], ['p5', 1, 'dog'], ['p6', 1, 'cat'], ['p7', 1, 'dog'], ['p8', 2, 'cat'],['p3', 2, 'cat'], ['p4', 2, 'cat'], ['p12', 3, 'cat'], ['p9', 3, 'cat'], ['p10', 3, 'dog'], ['p11', 3, 'dog'], ['p13', 4, 'cat']] sdf = spark.createDataFrame(data, schema = ['id', 'rank', 'group']) sdf.show()
输出:
+---+----+-----+ | id|rank|group| +---+----+-----+ | p1| 0| dog| | p2| 0| dog| | p5| 1| dog| | p6| 1| cat| | p7| 1| dog| | p8| 2| cat| | p3| 2| cat| | p4| 2| cat| |p12| 3| cat| | p9| 3| cat| |p10| 3| dog| |p11| 3| dog| |p13| 4| cat| +---+----+-----+
实现代码
方法:窗口函数计算分组特征后过滤
from pyspark.sql import Window import pyspark.sql.functions as F # 定义窗口:按rank分组 rank_window = Window.partitionBy("rank") # 计算每个rank分组的distinct group数量、分组行数 sdf_with_stats = sdf.withColumn( "distinct_group_count", F.countDistinct("group").over(rank_window) ).withColumn( "group_size", F.count("*").over(rank_window) ) # 过滤条件:distinct_group_count>1 或者 (distinct_group_count==1且group_size==1) filtered_sdf = sdf_with_stats.filter( (F.col("distinct_group_count") > 1) | ((F.col("distinct_group_count") == 1) & (F.col("group_size") == 1)) ).drop("distinct_group_count", "group_size") filtered_sdf.show()
输出结果
+---+----+-----+ | id|rank|group| +---+----+-----+ | p5| 1| dog| | p6| 1| cat| | p7| 1| dog| |p12| 3| cat| | p9| 3| cat| |p10| 3| dog| |p11| 3| dog| |p13| 4| cat| +---+----+-----+
逻辑说明
- 通过窗口函数
partitionBy("rank"),为每一行计算所属rank分组的不同group值的数量和分组总行数 - 过滤时保留两类分组:
- 分组内有多种group值(
distinct_group_count>1) - 分组仅一行(即使group值唯一,
group_size==1)
- 分组内有多种group值(
- 最后移除辅助计算的列,得到目标结果
内容的提问来源于stack exchange,提问作者Rory
相关产品推荐
相关产品推荐

