PySpark DataFrame基于非空值的重叠窗口collect_list实现
PySpark 按条件收集指定行数的列值到数组
给定已按name和date升序排序的PySpark DataFrame:
+----+----------+----+---+ |name| date| v1| v2| +----+----------+----+---+ | a|2000-01-01|null|0.1| | a|2000-01-02| 1|0.2| | a|2000-01-03|null|0.3| | a|2000-01-04|null|0.4| | a|2000-01-05| 2|0.5| | a|2000-01-06|null|0.6| | a|2000-01-07|null|0.7| | b|2000-01-08|null|0.8| +----+----------+----+---+
需求:针对v1列的每个非空条目,从当前行开始收集v2列的x条记录到新列acc(窗口允许重叠)。当x=4时,期望输出:
+----+----------+----+---+--------------------+ |name| date| v1| v2| acc| +----+----------+----+---+--------------------+ | a|2000-01-01|null|0.1| null| | a|2000-01-02| 1|0.2|[0.2, 0.3, 0.4, 0.5]| | a|2000-01-03|null|0.3| null| | a|2000-01-04|null|0.4| null| | a|2000-01-05| 2|0.5|[0.5, 0.6, 0.7] | | a|2000-01-06|null|0.6| null| | a|2000-01-07|null|0.7| null| | b|2000-01-08|null|0.8| null| +----+----------+----+---+--------------------+
解决方案
通过窗口函数结合条件判断就能实现需求,具体代码如下(以x=4为例):
from pyspark.sql import SparkSession from pyspark.sql import Window import pyspark.sql.functions as F # 初始化SparkSession spark = SparkSession.builder.appName("collect_v2").getOrCreate() # 构建示例DataFrame data = [ ("a", "2000-01-01", None, 0.1), ("a", "2000-01-02", 1, 0.2), ("a", "2000-01-03", None, 0.3), ("a", "2000-01-04", None, 0.4), ("a", "2000-01-05", 2, 0.5), ("a", "2000-01-06", None, 0.6), ("a", "2000-01-07", None, 0.7), ("b", "2000-01-08", None, 0.8) ] df = spark.createDataFrame(data, ["name", "date", "v1", "v2"]) x = 4 # 定义窗口:按name分区,从当前行到后续x-1行(包含当前行) window_spec = Window.partitionBy("name").orderBy("date").rowsBetween(0, x-1) # 生成acc列:仅在v1非空时保留收集结果,其余行设为null result_df = df.withColumn( "acc", F.when( F.col("v1").isNotNull(), F.collect_list("v2").over(window_spec) ).otherwise(F.lit(None)) ) # 查看结果 result_df.show(truncate=False)
关键逻辑说明
rowsBetween(0, x-1):指定窗口范围从当前行开始,到后续x-1行结束,保证最多收集包括当前行在内的x条v2值F.when(...):仅对v1非空的行生成acc数组,其余行直接设为null,符合需求- 由于原DataFrame已按
name和date排序,窗口内的v2值会自动按时间顺序收集
内容的提问来源于stack exchange,提问作者user13877981
相关产品推荐
相关产品推荐

