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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 13:57:13