优化PySpark DataFrame含字典列的分组聚合性能
问题描述
我需要处理如下结构的PySpark DataFrame:
| name | countries | options |
|---|---|---|
| dummy1 | ["UK", "FR"] | [{"id": 1, "key1": 10, "key2": 20}, {"id": 2, "key1": 10, "key2": 20}] |
| dummy1 | ["FR"] | [{"id": 1, "key1": 20, "key2": 30}] |
| dummy2 | ["UK", "FR"] | [{"id": 1, "key1": 10, "key2": 20}] |
转换为如下结构:
| name | countries | options |
|---|---|---|
| dummy1 | ["UK", "FR"] | [{"id": 1, "key1": 20, "key2": 30}, {"id": 2, "key1": 10, "key2": 20}] |
| dummy2 | ["UK", "FR"] | [{"id": 1, "key1": 10, "key2": 20}] |
核心需求
- 按
name分组,得到去重的countries列表; - 对于
options,按name+id分组后保留key1值最大的项,若key1相同则保留key2最大的项。
当前实现及问题
目前我是分别聚合countries和options后再做join,但数据量较大时运行速度极慢。
聚合countries的代码
countries = df.groupBy(col("name")) \ .agg(array_distinct(flatten(collect_list(col("countries")))).alias("countries"))
聚合options的代码
window = Window().partitionBy("name", "id").orderBy(col("key1").desc(), col("key2").desc()) options = df.select("name", explode("options").alias("options")) \ .withColumn("id", col("options.id")) \ .withColumn("key1", col("options.key1")) \ .withColumn("key2", col("options.key2")) \ .withColumn("rn", row_number().over(window)) \ .filter("rn = 1") \ .groupBy(col("name")) \ .agg(collect_list("options").alias("options"))
性能优化方案
1. 合并聚合逻辑,减少Shuffle次数
当前方案做了两次groupBy和一次join,会产生多次Shuffle。可以将两个聚合逻辑合并,只做一次Shuffle:
from pyspark.sql import functions as F from pyspark.sql.window import Window # 展开options并保留countries数据 df_expanded = df.select( "name", "countries", F.explode("options").alias("option") ).withColumn("id", F.col("option.id")) \ .withColumn("key1", F.col("option.key1")) \ .withColumn("key2", F.col("option.key2")) # 为每个name+id组选出最优option window_opt = Window.partitionBy("name", "id").orderBy(F.col("key1").desc(), F.col("key2").desc()) df_opt = df_expanded.withColumn("rn", F.row_number().over(window_opt)) \ .filter(F.col("rn") == 1) \ .drop("rn", "id", "key1", "key2") # 一次groupBy完成两个聚合操作 result = df_opt.groupBy("name") \ .agg( F.array_distinct(F.flatten(F.collect_list("countries"))).alias("countries"), F.collect_list("option").alias("options") )
2. 提前分区优化Window函数
对数据提前按name+id分区,减少Window函数执行时的Shuffle开销:
from pyspark.sql import functions as F from pyspark.sql.window import Window df_expanded = df.select( "name", "countries", F.explode("options").alias("option") ).withColumn("id", F.col("option.id")) \ .withColumn("key1", F.col("option.key1")) \ .withColumn("key2", F.col("option.key2")) \ .repartition(F.col("name"), F.col("id")) # 提前按分组键分区 window_opt = Window.partitionBy("name", "id").orderBy(F.col("key1").desc(), F.col("key2").desc()) df_opt = df_expanded.withColumn("rn", F.row_number().over(window_opt)) \ .filter(F.col("rn") == 1) \ .drop("rn", "id", "key1", "key2") result = df_opt.groupBy("name") \ .agg( F.array_distinct(F.flatten(F.collect_list("countries"))).alias("countries"), F.collect_list("option").alias("options") )
3. 优化countries聚合函数
用aggregate+array_union替代array_distinct(flatten(...)),避免先展平再去重的额外开销:
# 替换原countries聚合逻辑 F.aggregate( F.collect_list("countries"), F.array().cast("array<string>"), lambda acc, x: F.array_union(acc, x) ).alias("countries")
4. 基础数据预处理优化
- 调整DataFrame分区数:根据集群资源和数据量,用
df = df.repartition(n)设置合理的分区数; - 提前过滤无用数据:删除不需要的行或列,减少后续处理的数据量。
5. 开启Spark自适应执行
在Spark配置中设置spark.sql.adaptive.enabled=true,让Spark自动优化执行计划,比如合并Shuffle分区、调整Join策略等。
内容的提问来源于stack exchange,提问作者Francisco Fonseca
相关产品推荐
相关产品推荐

