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

Spark基于列值动态移位行:无需UDF的分组移位实现问询

无需UDF实现Spark按组移位提取日期

输入DataFrame

+---+----------+--------+
|ID |date      |shift_by|
+---+----------+--------+
|1  |2021-01-01|2       |
|1  |2021-02-05|2       |
|1  |2021-03-27|2       |
|2  |2022-02-28|1       |
|2  |2022-04-30|1       |
+---+----------+--------+

需求

按ID分组,基于shift_by的值执行移位操作:

  • 取组内第一条日期作为date1
  • 取组内相对于第一条偏移shift_by位的日期作为date2(如ID=1的shift_by=2,取第3条日期;ID=2的shift_by=1,取第2条日期)

最终结果如下:

+---+----------+----------+
|ID |date1     |date2     |
+---+----------+----------+
|1  |2021-01-01|2021-03-27|
|2  |2022-02-28|2022-04-30|
+---+----------+----------+

当前问题

已通过UDF实现逻辑,但运行效率低下,寻求不依赖UDF的优化方案。

解决方案

方案一:分组聚合+内置函数(推荐)

利用Spark内置聚合函数直接处理分组数据,避免UDF的性能开销:

首先创建输入DataFrame:

from datetime import datetime
from pyspark.sql import SparkSession
from pyspark.sql.types import *

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

input_data = [
    (1, datetime.date(2021, 1, 1), 2),
    (1, datetime.date(2021, 2, 5), 2),
    (1, datetime.date(2021, 3, 27), 2),
    (2, datetime.date(2022, 2, 28), 1),
    (2, datetime.date(2022, 4, 30), 1)
]
input_schema = StructType([
    StructField("ID", IntegerType(), True),
    StructField("date", DateType(), True),
    StructField("shift_by", IntegerType(), True)
])
input_df = spark.createDataFrame(data=input_data, schema=input_schema)

执行聚合逻辑:

from pyspark.sql.functions import collect_list, first, element_at, sort_array

# 按ID分组,收集排序后的日期列表,取第一个元素和偏移shift_by后的元素
result_df = input_df.groupBy("ID")\
    .agg(
        first("date").alias("date1"),
        # element_at索引从1开始,因此取shift_by+1位置的元素
        element_at(sort_array(collect_list("date")), first("shift_by") + 1).alias("date2")
    )

result_df.show()

方案二:窗口函数实现

通过窗口函数为组内行编号,筛选目标行后转为宽表:

from pyspark.sql.window import Window
from pyspark.sql.functions import row_number, max, when, first

# 定义窗口:按ID分组,日期排序
window = Window.partitionBy("ID").orderBy("date")

# 添加组内行号
ranked_df = input_df.withColumn("row_num", row_number().over(window))

# 获取每个ID对应的目标行号(shift_by+1)
target_row_df = ranked_df.groupBy("ID")\
    .agg(max(when(ranked_df.row_num == 1, ranked_df.shift_by + 1)).alias("target_row"))

# 筛选行号为1和目标行号的记录
filtered_df = ranked_df.join(target_row_df, on="ID")\
    .filter((ranked_df.row_num == 1) | (ranked_df.row_num == target_row_df.target_row))

# 转宽表得到结果
result_df = filtered_df.groupBy("ID")\
    .pivot("row_num")\
    .agg(first("date"))\
    .withColumnRenamed("1", "date1")\
    .withColumnRenamed(str(filtered_df.select("target_row").first()[0]), "date2")

result_df.show()

说明

方案一性能更优,仅需一次分组聚合操作,避免了窗口函数带来的额外Shuffle开销,适合同一ID的shift_by值统一的场景(如输入数据所示)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 02:01:12