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

基于条件的Spark窗口函数:计算Cond=1的前两日Val值之和

问题描述

现有如下Spark DataFrame:

|-----------------------|
|Date       | Val | Cond|
|-----------------------|
|2022-01-08 | 2   | 0   |
|2022-01-09 | 4   | 1   |
|2022-01-10 | 6   | 1   |
|2022-01-11 | 8   | 0   |
|2022-01-12 | 2   | 1   |
|2022-01-13 | 5   | 1   |
|2022-01-14 | 7   | 0   |
|2022-01-15 | 9   | 0   | 
|-----------------------|

需要为每个日期计算其之前所有满足Cond=1的最近2个Val值之和,预期输出如下:

|-----------------|
|Date       | Sum |
|-----------------|
|2022-01-08 | 0   |  不计算,因为该日期前没有两个Cond=1的日期
|2022-01-09 | 0   |  不计算,因为该日期前没有两个Cond=1的日期
|2022-01-10 | 0   |  不计算,因为该日期前没有两个Cond=1的日期
|2022-01-11 | 10  | (4+6)
|2022-01-12 | 10  | (4+6)
|2022-01-13 | 8   | (2+6)
|2022-01-14 | 7   | (5+2)
|2022-01-15 | 7   | (5+2)
|-----------------|

尝试了以下代码,但使用.where("Cond=1")会直接排除Cond=0的日期,导致结果缺失这些行:

df = df.where("Cond= 1").withColumn(
    "ListView",
    f.collect_list("Val").over(windowSpec.rowsBetween(-2, -1))
)

如何用窗口函数实现预期输出?

最小可复现代码(MVCE):

from pyspark.sql import SparkSession
from pyspark.sql.types import StructType, StructField, DateType, IntegerType
import pyspark.sql.functions as f

spark = SparkSession.builder.appName("example").getOrCreate()

data_1=[
    ("2022-01-08",2,0),
    ("2022-01-09",4,1),
    ("2022-01-10",6,1),
    ("2022-01-11",8,0),
    ("2022-01-12",2,1),
    ("2022-01-13",5,1),
    ("2022-01-14",7,0),
    ("2022-01-15",9,0) 
]

schema_1 = StructType([
    StructField("Date", DateType(),True),
    StructField("Val", IntegerType(),True),
    StructField("Cond", IntegerType(),True)
  ])

df_1 = spark.createDataFrame(data=data_1,schema=schema_1)
解决方案

核心思路是不提前过滤行,而是在窗口函数内部对Cond=1的Val值进行收集,再取最近2个求和:

  1. 定义全局窗口(按Date排序,范围是从起始到当前行的前一行)
  2. 用collect_list收集所有满足Cond=1的Val值,过滤无效值后取最近的2个元素求和
  3. 处理不足2个元素的情况,返回0

代码实现:

from pyspark.sql.window import Window

# 定义窗口:按Date升序,覆盖当前日期之前的所有行
window_spec = Window.orderBy("Date").rowsBetween(Window.unboundedPreceding, -1)

# 计算目标Sum列
result_df = df_1.withColumn(
    # 收集所有Cond=1的Val,其余设为null
    "valid_vals",
    f.collect_list(f.when(f.col("Cond") == 1, f.col("Val"))).over(window_spec)
).withColumn(
    # 过滤列表中的null值,得到仅符合条件的Val列表
    "filtered_vals",
    f.expr("filter(valid_vals, x -> x is not null)")
).withColumn(
    # 反转列表取最近2个元素求和,不足2个则返回0
    "Sum",
    f.when(
        f.size(f.col("filtered_vals")) >= 2,
        f.expr("aggregate(slice(reverse(filtered_vals), 1, 2), 0, (acc, x) -> acc + x)")
    ).otherwise(0)
).select("Date", "Sum")

result_df.show()

代码说明

  • window_spec:确保能获取当前日期之前的所有历史行,不会遗漏符合条件的记录
  • valid_vals:保留Cond=1的Val,其余标记为null,避免无意义值混入
  • filtered_vals:清理null值,得到纯符合条件的Val列表
  • Sum:通过反转列表+切片获取最近2个符合条件的Val,用aggregate求和;列表长度不足2时直接返回0

执行后输出与预期完全一致:

+----------+---+
|      Date|Sum|
+----------+---+
|2022-01-08|  0|
|2022-01-09|  0|
|2022-01-10|  0|
|2022-01-11| 10|
|2022-01-12| 10|
|2022-01-13|  8|
|2022-01-14|  7|
|2022-01-15|  7|
+----------+---+

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 21:25:21