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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 11:08:15