PySpark基于条件的分区累计求和实现求助
PySpark实现条件累加cumulative_pass列
需求描述
针对每个username分区,生成cumulative_pass列,该列的值为所有满足当前行year_start大于之前行year_end的历史行的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| +--------+-------+----------+--------+----+---------------+
原尝试代码的问题
原代码存在以下错误:
- 未导入
IntegerType,需要从pyspark.sql.types中导入 - Pandas UDF逻辑错误:窗口内的
df['year_start'].max()是窗口所有行的year_start最大值,而非当前行的year_start - 窗口函数与UDF的结合方式不正确,无法正确传递当前行的判断条件
正确实现方法
无需使用UDF,直接通过窗口函数结合条件判断即可高效实现:
from pyspark.sql.types import IntegerType # 定义窗口:按username分区,按session排序,取当前行之前的所有历史行 w = Window.partitionBy("username") \ .orderBy("session") \ .rowsBetween(Window.unboundedPreceding, Window.currentRow - 1) # 计算cumulative_pass:累加历史行中year_end < 当前行year_start的pass值 df_result = df.withColumn( "cumulative_pass", F.sum( F.when(F.col("year_end") < F.col("year_start"), F.col("pass")) .otherwise(0) ).over(w) ) df_result.show()
代码说明
- 窗口
w限定了每个username下,仅对当前行之前的所有行进行计算 - 使用
F.when过滤出满足year_end < 当前行year_start的历史行,对这些行的pass求和,不满足条件的行贡献0 - 最终生成的
cumulative_pass完全符合需求
内容的提问来源于stack exchange,提问作者s_i
相关产品推荐
相关产品推荐

