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

将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()

逻辑说明

  1. 窗口定义:与原UDF使用的窗口范围一致,确保计算范围覆盖当前用户的历史行到当前行。
  2. 当前行year_start获取:利用窗口内的max函数直接拿到当前行的year_start,等价于原UDF中df['year_start'].max()的逻辑。
  3. 有效pass筛选:通过when函数完成条件判断,避免了UDF的序列化开销。
  4. 累加求和:用PySpark原生的窗口求和函数完成累加,性能远高于Pandas UDF。

这个实现完全基于PySpark原生API,消除了UDF带来的性能瓶颈,处理大数据量时优势会更明显。


内容的提问来源于stack exchange,提问作者s_i

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 16:15:08