将Pandas UDF改写为纯PySpark窗口函数以优化性能
用纯PySpark窗口函数替代Pandas UDF优化性能生成
cumulative_pass列 我需要将现有代码中的Pandas UDF替换为纯PySpark窗口函数,以此优化性能,目标是程序化生成cumulative_pass列。
初始数据与代码
import pandas as pd import pyspark.sql.functions as F from pyspark.sql import SparkSession from pyspark.sql import Window import sys spark_session = SparkSession.builder.getOrCreate() df_data = {'username': ['bob','bob', 'bob', 'bob', 'bob', 'bob', 'bob', 'bob'], 'session': [1,2,3,4,5,6,7,8], 'year_start': [2020,2020,2020,2020,2020,2021,2022,2023], 'year_end': [2020,2020,2020,2020,2021,2021,2022,2023], 'pass': [1,0,0,0,0,1,1,0], 'cumulative_pass': [0,0,0,0,0,1,2,3], } df_pandas = pd.DataFrame.from_dict(df_data) df = spark_session.createDataFrame(df_pandas) df.show()
执行后输出:
+--------+-------+----------+--------+----+---------------+ |username|session|year_start|year_end|pass|cumulative_pass| +--------+-------+----------+--------+----+---------------+ | bob| 1| 2020| 2020| 1| 0| | bob| 2| 2020| 2020| 0| 0| | bob| 3| 2020| 2020| 0| 0| | bob| 4| 2020| 2020| 0| 0| | bob| 5| 2020| 2021| 0| 0| | bob| 6| 2021| 2021| 1| 1| | bob| 7| 2022| 2022| 1| 2| | bob| 8| 2023| 2023| 0| 3| +--------+-------+----------+--------+----+---------------+
当前性能较慢的Pandas UDF实现
def conditional_sum(data: pd.DataFrame) -> int: df = data.apply(pd.Series) return df.loc[df['year_start'].max() > df['year_end']]['pass'].sum() udf_conditional_sum = F.pandas_udf(conditional_sum, IntegerType()) w = Window.partitionBy("username").orderBy(F.asc("year_start")).rowsBetween(-sys.maxsize, 0) df = df.withColumn("calculate_cumulative_pass", udf_conditional_sum(F.struct("year_start", "year_end", "pass")).over(w))
注:已对窗口w做了小幅修改,并移除了二次排序。
纯PySpark窗口函数解决方案
观察cumulative_pass的生成逻辑:对每个用户,按year_start升序排列后,当前行的cumulative_pass等于所有之前行中满足当前行year_start > 该行year_end的pass值之和。
用纯PySpark窗口函数可以这样实现:
from pyspark.sql.types import IntegerType # 定义窗口:按username分区,按year_start升序,覆盖从起始行到当前行的范围 window_spec = Window.partitionBy("username").orderBy(F.asc("year_start")).rowsBetween(Window.unboundedPreceding, Window.currentRow) # 获取当前窗口内的year_start最大值(即当前行的year_start) current_year_start = F.max("year_start").over(window_spec) # 标记符合条件的pass值:当前行year_start > 该行year_end时保留pass,否则为0 valid_pass = F.when(current_year_start > F.col("year_end"), F.col("pass")).otherwise(0) # 累加有效pass值得到目标列 df = df.withColumn("calculate_cumulative_pass", F.sum(valid_pass).over(window_spec)) df.show()
逻辑说明
- 窗口定义:与原UDF使用的窗口范围一致,确保计算范围覆盖当前用户的历史行到当前行。
- 当前行year_start获取:利用窗口内的
max函数直接拿到当前行的year_start,等价于原UDF中df['year_start'].max()的逻辑。 - 有效pass筛选:通过
when函数完成条件判断,避免了UDF的序列化开销。 - 累加求和:用PySpark原生的窗口求和函数完成累加,性能远高于Pandas UDF。
这个实现完全基于PySpark原生API,消除了UDF带来的性能瓶颈,处理大数据量时优势会更明显。
内容的提问来源于stack exchange,提问作者s_i
相关产品推荐
相关产品推荐

