PySpark中合并多大数据DataFrame的最优方案及persist用法咨询
我有N个(100-1000个)超大DataFrame(单个大小100GB-1TB),需对每个DataFrame处理得到数MB级小DataFrame后合并。当前采用循环union方式实现(代码如下),请问是否有更优方案?另外,循环中是否需用.persist()提速?
df_merged = new_empty_df for i in range(N): df_raw = read(i) df_processed = process(df_raw) df_merged = df_merged.union(df_processed) # should .persist() be used here? # output processed df_merged.write(somewhere)
补充说明:数据为含GPS位置的tsv/parquet文件,process函数先为每个Location点计算AreaId,再按AreaId统计点数,代码如下:
df_area_id= df_raw.withColumn('AreaId',GetTileForPointUDF(col('Location'))) df_processed = df_area_id.groupby('AreaId').count()
相关数据示例:
原始数据:
| Location |
|---|
| POINT (30 10) |
| POINT (33.44 -22.4) |
| POINT (33.11 -21.7) |
计算AreaId后的数据:
| Location | AreaId |
|---|---|
| POINT (30 10) | 2 |
| POINT (33.44 -22.4) | 4 |
| POINT (33.11 -21.7) | 4 |
统计结果:
| AreaId | Count |
|---|---|
| 2 | 1 |
| 4 | 2 |
一、别用循环Union了,换全局/分层聚合才是最优解
你当前的循环Union方案有个致命问题:每Union一次,Spark的执行计划就会膨胀一圈,当N到1000的时候,执行计划会复杂到让Spark优化器直接卡壳,调度开销也会爆炸。而且你的场景刚好是每个小DF都是按AreaId聚合后的统计结果,完全可以跳过“局部聚合→Union合并”的弯路,直接做全局聚合:
方案1:一次性读取全量数据做全局聚合
如果你的存储系统能支撑(比如HDFS、S3这类分布式存储),直接把所有输入文件当成一个数据源读取,然后在全量数据上执行处理逻辑:
# 直接读整个目录下的所有parquet/tsv文件,不用循环 df_raw = spark.read.format("parquet").load("/path/to/all/input_files/") # 或者传文件路径列表,效果一样 # file_paths = [f"/path/to/file_{i}" for i in range(N)] # df_raw = spark.read.format("parquet").load(file_paths) # 直接全局计算AreaId+聚合 df_area_id = df_raw.withColumn('AreaId', GetTileForPointUDF(col('Location'))) df_final = df_area_id.groupby('AreaId').count() df_final.write(somewhere)
这种方式让Spark自己优化执行计划,不管数据量多大,性能都比循环Union强10倍以上,代码还简洁。
方案2:分区聚合+全局二次聚合(无法全量读取时用)
如果因为文件权限、存储限制没法一次性读全量,那就把“Union合并”改成“累加聚合”——每次处理完一个大文件,就把局部聚合的结果和全局统计结果合并(相同AreaId的count直接相加),这样最终的DF始终是MB级,不会随着N变大而膨胀:
from pyspark.sql.types import StructType, StructField, IntegerType, LongType # 初始化全局聚合结果的空DF,指定schema agg_schema = StructType([ StructField("AreaId", IntegerType(), nullable=False), StructField("count", LongType(), nullable=False) ]) df_global_agg = spark.createDataFrame([], agg_schema) for i in range(N): df_raw = read(i) df_area_id = df_raw.withColumn('AreaId', GetTileForPointUDF(col('Location'))) # 先做局部文件的聚合 df_local_agg = df_area_id.groupby('AreaId').count() # 把局部结果和全局结果合并,直接按AreaId累加count df_global_agg = df_global_agg.union(df_local_agg) \ .groupBy('AreaId').sum('count') \ .withColumnRenamed('sum(count)', 'count') df_global_agg.write(somewhere)
这个方案的执行计划不会随着N变大而爆炸,而且内存占用始终很低,比循环Union靠谱多了。
二、循环里的.persist()?完全没必要,搞不好还拖后腿
别加.persist(),原因很简单:
- 你每次Union后都会生成新的
df_merged,旧的DF会被自动回收,加persist只会把越来越大的合并结果塞进内存/磁盘,白白占用资源,反而拖慢速度。 - 你的
df_processed是MB级的小数据,Spark处理这种小DF根本不需要缓存,瞬间就能完成合并。 - 要是用了上面的聚合方案,persist更是多余——局部聚合后的结果本来就小,全局聚合的逻辑会被Spark优化器自动处理,手动缓存纯属画蛇添足。
额外的性能优化 tips
- 把Python UDF换成Pandas UDF或者Scala UDF:Python UDF的序列化开销极大,TB级数据跑起来会慢到离谱。如果是Spark 3.x+,也可以试试用Spark内置的地理函数替代自定义UDF,性能提升明显。
- 调整Spark配置:针对超大文件,把
spark.sql.files.maxPartitionBytes调大(比如从默认128MB改成1GB),减少分区数量,降低调度开销;同时给executor和driver加够内存,避免OOM。 - 把TSV转成Parquet再处理:Parquet是列式存储,压缩比高,读取速度比TSV快好几倍,能省大量IO时间。
内容的提问来源于stack exchange,提问作者Alcibiades

