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
相关产品推荐
相关产品推荐

