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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 14:45:34