基于条件的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个求和:
- 定义全局窗口(按Date排序,范围是从起始到当前行的前一行)
- 用
collect_list收集所有满足Cond=1的Val值,过滤无效值后取最近的2个元素求和 - 处理不足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
相关产品推荐
相关产品推荐

