如何在PySpark中高效筛选满足组特定条件的DataFrame分组
如何高效排除不满足组依赖过滤条件的数据分组?
测试数据
我们使用以下测试数据:
df = spark.createDataFrame([(1,2),(1,3),(1,40),(1,0),(2,3),(2,1),(2,4),(3,2),(3,4)],['a','b']) df.show()
输出结果:
+---+---+ | a| b| +---+---+ | 1| 2| | 1| 3| | 1| 40| | 1| 0| | 2| 3| | 2| 1| | 2| 4| | 3| 2| | 3| 4| +---+---+
需求
- 过滤掉
average(b) ≤ 6的数据分组
预期输出
+---+---+ | a| b| +---+---+ | 1| 2| | 1| 3| | 1| 40| | 1| 0| +---+---+
当前实现及问题
当前实现代码:
df_filter = df.groupby('a').agg(F.mean(F.col('b')).alias("avg")) df_filter = df_filter.filter(F.col('avg') > 6.) df.join(df_filter,'a','inner').drop('avg').show()
存在的核心问题:
- 触发两次Shuffle:一次是分组计算均值时,另一次是执行join操作时。对应的物理执行计划如下:
== Physical Plan == *(5) Project [a#175L, b#176L] +- *(5) SortMergeJoin [a#175L], [a#222L], Inner :- *(2) Sort [a#175L ASC NULLS FIRST], false, 0 : +- Exchange hashpartitioning(a#175L, 200), ENSURE_REQUIREMENTS, [plan_id=919] : +- *(1) Filter isnotnull(a#175L) : +- *(1) Scan ExistingRDD[a#175L,b#176L] +- *(4) Sort [a#222L ASC NULLS FIRST], false, 0 +- *(4) Project [a#222L] +- *(4) Filter (isnotnull(avg#219) AND (avg#219 > 6.0)) +- *(4) HashAggregate(keys=[a#222L], functions=[avg(b#223L)]) +- Exchange hashpartitioning(a#222L, 200), ENSURE_REQUIREMENTS, [plan_id=925] +- *(3) HashAggregate(keys=[a#222L], functions=[partial_avg(b#223L)]) +- *(3) Filter isnotnull(a#222L) +- *(3) Scan ExistingRDD[a#222L,b#223L]
高效解决方案
方法一:窗口函数(最优,仅一次Shuffle)
通过窗口函数计算每个分组的均值,直接过滤满足条件的行,全程仅需一次Shuffle:
from pyspark.sql import Window import pyspark.sql.functions as F # 定义按a分组的窗口规则 window_spec = Window.partitionBy('a') # 计算分组均值并过滤 result_df = df.withColumn('avg_b', F.mean('b').over(window_spec)) \ .filter(F.col('avg_b') > 6.) \ .drop('avg_b') result_df.show()
验证执行计划:
result_df.explain()
计划中只会出现一次Exchange hashpartitioning(a#xxx, 200),即仅一次Shuffle操作——窗口函数计算时按a分区完成后,后续过滤无需再次Shuffle。
方法二:广播过滤后的分组Key(适用于分组数量少的场景)
如果过滤后保留的分组数量极少,可将符合条件的a值广播,避免join阶段的Shuffle:
# 获取符合条件的a值列表 valid_a = df.groupby('a').agg(F.mean('b').alias('avg')) \ .filter(F.col('avg') > 6.) \ .select('a') \ .rdd.flatMap(lambda x: x).collect() # 基于广播的小数据集过滤原表 result_df = df.filter(F.col('a').isin(valid_a)) result_df.show()
这种方式仅在分组计算均值时触发一次Shuffle,后续过滤依赖广播的小数据集,无额外Shuffle。但如果分组数量较多,广播会占用过多内存,此时窗口函数是更优选择。
内容的提问来源于stack exchange,提问作者figs_and_nuts
相关产品推荐
相关产品推荐

