PySpark DataFrame按条件获取下一个工作日日期的实现问题
问题描述
给定如下PySpark DataFrame:
| date | is_business_day |
|---|---|
| 2023-01-01 | 0 |
| 2023-01-02 | 1 |
| 2023-01-03 | 1 |
| 2023-01-04 | 1 |
| 2023-01-05 | 1 |
| 2023-01-06 | 1 |
| 2023-01-07 | 0 |
| 2023-01-08 | 0 |
| 2023-01-09 | 1 |
| 2023-04-06 | 1 |
| 2023-04-07 | 0 |
| 2023-04-08 | 0 |
| 2023-04-09 | 0 |
| 2023-04-10 | 1 |
需要为每行添加next_business_day列,值为该行日期之后第一个满足is_business_day == 1的日期。
自行编写的next_business_day函数无法在.withColumn()中使用,执行以下代码时:
df_calendar = ( df_calendar .withColumn('next_business_day', next_business_day(df_calendar, col('date'))) )
抛出错误:
TypeError: Invalid argument, not a string or column: DataFrame[date: date] of type <class 'pyspark.sql.dataframe.DataFrame'>. For column literals, use 'lit', 'array', 'struct' or 'create_map' function.
问题原因
.withColumn()的第二个参数要求是列表达式,不能直接传入Python函数调用(除非是UDF),且UDF无法接收整个DataFrame作为参数。- 原函数中使用
collect()会将全量数据拉取到Driver端,不仅效率极低,还可能导致内存溢出,完全违背Spark的分布式计算设计。
解决方案
推荐两种高效的分布式实现方式:
方法1:广播工作日映射表
利用工作日的映射关系,通过广播小数据集来实现高效关联:
from pyspark.sql import functions as F from pyspark.sql.window import Window # 确保date列是date类型(若原数据是字符串需转换) df_calendar = df_calendar.withColumn("date", F.to_date("date")) # 提取所有工作日,并生成每个工作日的下一个工作日 business_days = df_calendar.filter(F.col("is_business_day") == 1).select("date").orderBy("date") business_days_with_next = business_days.withColumn( "next_business_day", F.lead("date").over(Window.orderBy("date")) ) # 广播映射表(因为工作日数据量远小于全量日历,广播后避免shuffle) broadcast_bd = F.broadcast(business_days_with_next) # 关联并找到每个日期之后的第一个工作日 df_result = df_calendar.join( broadcast_bd, df_calendar["date"] < broadcast_bd["date"], "left" ).groupBy(df_calendar["date"], df_calendar["is_business_day"]).agg( F.min(broadcast_bd["date"]).alias("next_business_day") )
方法2:窗口函数+向前填充
通过窗口函数标记并填充后续的第一个工作日:
from pyspark.sql import functions as F from pyspark.sql.window import Window # 确保date列是date类型 df_calendar = df_calendar.withColumn("date", F.to_date("date")) # 创建按日期升序的窗口 window_spec = Window.orderBy("date") # 标记工作日的日期,然后用last函数取后续第一个非空的工作日日期 df_result = df_calendar.withColumn( # 仅工作日保留自身日期,非工作日设为null "business_day_mark", F.when(F.col("is_business_day") == 1, F.col("date")) ).withColumn( # 从当前行的下一行开始,取最近的非空工作日日期 "next_business_day", F.last("business_day_mark", ignorenulls=True).over(window_spec.rowsBetween(1, Window.unboundedFollowing)) ).drop("business_day_mark")
验证结果
两种方法均可生成符合预期的输出:
| date | is_business_day | next_business_day |
|---|---|---|
| 2023-01-01 | 0 | 2023-01-02 |
| 2023-01-02 | 1 | 2023-01-03 |
| 2023-01-03 | 1 | 2023-01-04 |
| 2023-01-04 | 1 | 2023-01-05 |
| 2023-01-05 | 1 | 2023-01-06 |
| 2023-01-06 | 1 | 2023-01-09 |
| 2023-01-07 | 0 | 2023-01-09 |
| 2023-01-08 | 0 | 2023-01-09 |
| 2023-01-09 | 1 | 2023-04-06 |
| 2023-04-06 | 1 | 2023-04-10 |
| 2023-04-07 | 0 | 2023-04-10 |
| 2023-04-08 | 0 | 2023-04-10 |
| 2023-04-09 | 0 | 2023-04-10 |
| 2023-04-10 | 1 | null |
注:最后一行的
next_business_day为null,因为没有后续的工作日数据,可根据需求用F.fillna处理。
内容的提问来源于stack exchange,提问作者Marcos Martins
相关产品推荐
相关产品推荐

