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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 13:49:53