Spark循环内操作DataFrame每次迭代运行速度越来越慢问题
问题背景
- Spark 初学者在实现多次遍历 DataFrame 的需求时遇到性能问题:对 DataFrame 执行10次带不同日期参数的筛选循环,随着循环次数增加,处理耗时线性上涨,手动调用
unpersist()清理缓存未取得优化效果。
问题复现代码如下:
import findspark findspark.init() import pyspark.sql.functions as F from pyspark.sql import SparkSession from itertools import combinations import datetime spark = SparkSession.builder.appName("Practice").master("local[*]").config("spark.executor.memory", "70g").config("spark.driver.memory", "50g").config("spark.memory.offHeap.enabled",True).config("spark.memory.offHeap.size","16g").getOrCreate() df = spark.read.parquet('spark-big-data\parquet_small_example.parquet') res =[] for date in range(10): df = df.withColumn('fs_origin',df.request.Segments.getItem(0)['Origin']) df = df.withColumn('fs_destination',df.request.Segments.getItem(0)['Destination']) df = df.withColumn('fs_date',df.request.Segments.getItem(0)['FlightTime']) df = df.withColumn('ss_origin',df.request.Segments.getItem(1)['Origin']) df = df.withColumn('ss_destination',df.request.Segments.getItem(1)['Destination']) df = df.withColumn('ss_date',df.request.Segments.getItem(1)['FlightTime']) df = df.withColumn('full_date',F.concat_ws('-', df.year,df.month,df.day)) df = df.filter( (df["fs_origin"] == 'TLV') & (df["fs_destination"] == 'NYC') & (df["ss_origin"] == 'NYC') & (df['ss_destination']=='TLV') & (df['fs_date']=='2021-02-'+str(date)+'T00:00:00') & (df['ss_date']=='2021-02-16'+'T00:00:00')) if df.count()==0: res.append(0) else: df = df.sort(F.unix_timestamp("full_date", "yyyy-M-d").desc()) latest_day = df.collect()[0]['full_date'] df = df.filter(df['full_date']==latest_day) df = df.withColumn("exploded_data", F.explode("response.results")) df = df.withColumn( "price", F.col("exploded_data").getItem('PriceInfo').getItem('Price') # either get by name or by index e.g. getItem(0) etc ) res.append(df.sort(df.price.asc()).collect()[0]['price']) df.unpersist() spark.catalog.clearCache()
核心问题原因
- 血缘无限膨胀:循环内每次都将处理结果赋值给原始
df变量,下一次循环的计算基于上一次循环的输出DataFrame展开,导致第N次循环的执行计划要包含前面N-1次所有的列计算、过滤逻辑,DAG长度随循环次数线性增长,计算开销越来越大。 - 重复计算冗余:航段信息提取、日期拼接这类固定逻辑在10次循环里重复执行,完全没有复用计算结果。
- 缓存使用错误:代码中从未对DataFrame执行
cache()/persist()操作,调用unpersist()和清空缓存的操作没有作用对象,属于无效操作;反而反复清空缓存会阻断Spark的自动优化能力。 - Action重复触发:单次循环内先调用
count()、再多次调用collect(),同一份数据被重复计算多次,额外增加开销。
优化实现方案
方案1:保留循环逻辑的修正写法
把固定不变的列派生逻辑全部提到循环外执行一次,处理完成后对基础数据集做缓存,循环内每次基于缓存好的基础DataFrame做筛选,不要迭代修改基础DataFrame变量:
import findspark findspark.init() import pyspark.sql.functions as F from pyspark.sql import SparkSession import datetime spark = SparkSession.builder.appName("Practice").master("local[*]")\ .config("spark.executor.memory", "70g")\ .config("spark.driver.memory", "50g")\ .config("spark.memory.offHeap.enabled",True)\ .config("spark.memory.offHeap.size","16g")\ .getOrCreate() # 所有固定列计算全部放到循环外,只执行一次 base_df = spark.read.parquet('spark-big-data\parquet_small_example.parquet')\ .withColumn('fs_origin', F.col('request.Segments').getItem(0)['Origin'])\ .withColumn('fs_destination', F.col('request.Segments').getItem(0)['Destination'])\ .withColumn('fs_date', F.col('request.Segments').getItem(0)['FlightTime'])\ .withColumn('ss_origin', F.col('request.Segments').getItem(1)['Origin'])\ .withColumn('ss_destination', F.col('request.Segments').getItem(1)['Destination'])\ .withColumn('ss_date', F.col('request.Segments').getItem(1)['FlightTime'])\ .withColumn('full_date', F.concat_ws('-', F.col('year'), F.col('month'), F.col('day')))\ .filter( (F.col("fs_origin") == 'TLV') & (F.col("fs_destination") == 'NYC') & (F.col("ss_origin") == 'NYC') & (F.col("ss_destination") == 'TLV') & (F.col("ss_date") == '2021-02-16T00:00:00') ) # 缓存基础数据集,供后续循环复用 base_df.cache() # 首次触发action完成缓存 base_df.count() res = [] for date in range(10): target_fs_date = f'2021-02-{date}T00:00:00' # 每次循环基于缓存好的base_df筛选,不要修改base_df本身 current_df = base_df.filter(F.col('fs_date') == target_fs_date) if current_df.count() == 0: res.append(0) continue # 取当前日期下最新full_date的最低价格 latest_ts = current_df.agg(F.max(F.unix_timestamp("full_date", "yyyy-M-d")).alias("max_ts")).collect()[0]['max_ts'] min_price = current_df.filter(F.unix_timestamp("full_date", "yyyy-M-d") == latest_ts)\ .select(F.explode("response.results").alias("exploded_data"))\ .select(F.col("exploded_data.PriceInfo.Price").alias("price"))\ .agg(F.min("price").alias("min_price"))\ .collect()[0]['min_price'] res.append(min_price) # 所有计算完成后再释放缓存 base_df.unpersist()
方案2:无循环的最优写法(推荐)
Spark本身是面向分布式集合的计算引擎,逐行循环的写法没有利用分布式计算的优势,可以直接通过分组聚合一次性算出所有目标日期的结果,完全消除循环开销,性能比循环写法高一个量级:
import findspark findspark.init() import pyspark.sql.functions as F from pyspark.sql import SparkSession from pyspark.sql.window import Window spark = SparkSession.builder.appName("Practice").master("local[*]")\ .config("spark.executor.memory", "70g")\ .config("spark.driver.memory", "50g")\ .config("spark.memory.offHeap.enabled",True)\ .config("spark.memory.offHeap.size","16g")\ .getOrCreate() # 生成要查询的所有目标日期列表 target_dates = [f'2021-02-{d}T00:00:00' for d in range(10)] result_df = spark.read.parquet('spark-big-data\parquet_small_example.parquet')\ .withColumn('fs_origin', F.col('request.Segments').getItem(0)['Origin'])\ .withColumn('fs_destination', F.col('request.Segments').getItem(0)['Destination'])\ .withColumn('fs_date', F.col('request.Segments').getItem(0)['FlightTime'])\ .withColumn('ss_origin', F.col('request.Segments').getItem(1)['Origin'])\ .withColumn('ss_destination', F.col('request.Segments').getItem(1)['Destination'])\ .withColumn('ss_date', F.col('request.Segments').getItem(1)['FlightTime'])\ .withColumn('full_date', F.concat_ws('-', F.col('year'), F.col('month'), F.col('day')))\ .filter( (F.col("fs_origin") == 'TLV') & (F.col("fs_destination") == 'NYC') & (F.col("ss_origin") == 'NYC') & (F.col("ss_destination") == 'TLV') & (F.col("ss_date") == '2021-02-16T00:00:00') & (F.col("fs_date").isin(target_dates)) )\ .withColumn("date_ts", F.unix_timestamp("full_date", "yyyy-M-d"))\ .withColumn("max_ts_per_fsdate", F.max("date_ts").over(Window.partitionBy("fs_date")))\ .filter(F.col("date_ts") == F.col("max_ts_per_fsdate"))\ .select(F.col("fs_date"), F.explode("response.results").alias("exploded_data"))\ .select(F.col("fs_date"), F.col("exploded_data.PriceInfo.Price").alias("price"))\ .groupBy("fs_date")\ .agg(F.min("price").alias("min_price")) # 收集结果,不存在的日期补0 price_map = {row['fs_date']: row['min_price'] for row in result_df.collect()} res = [price_map.get(d, 0) for d in target_dates]
优化要点总结
- 避免在循环中迭代赋值同一个DataFrame变量,防止执行计划血缘无限变长。
- 所有公共计算逻辑尽量前置,只计算一次后缓存复用,不要在循环中重复做相同的列转换。
- 缓存操作要先执行持久化方法,再通过action触发缓存生效,不要对未持久化的DataFrame执行
unpersist操作。 - 尽量用Spark原生的分组、窗口、集合操作替代Python层面的循环,减少Driver和Executor之间的交互次数,充分利用分布式计算能力。
内容的提问来源于stack exchange,提问作者Daniel Avigdor
相关产品推荐
相关产品推荐

