如何用Window函数替代GroupBy+Join实现Spark数据经纬度修正?
问题描述
我有一个结构为Grid_ID, Latitude, Longitude, DateTimeStamp的Spark DataFrame(df),输入数据如下:
Grid_ID Latitude Longitude DateTimeStamp Grid_1 Lat1 Long1 2021-06-30 00:00:00 Grid_1 Lat1 Long1 2021-06-30 00:01:00 Grid_1 Lat1 Long1 2021-06-30 00:02:00 Grid_1 Lat2 Long2 2021-07-01 00:00:00 Grid_1 Lat2 Long2 2021-07-01 00:01:00 Grid_1 Lat2 Long2 2021-07-01 00:02:00
已知同一Grid_ID对应的经纬度随日期互斥:日期≤2021-06-30时用Lat1/Long1,日期>2021-06-30时用Lat2/Long2。现在需要新增Corrected_Lat和Corrected_Long列,为所有行统一赋值Lat2/Long2。
目前我通过GroupBy+Agg+Join的方式实现,代码如下:
import pyspark.sql.functions as F df_dated = df.withColumn("date", F.to_date("DateTimeStamp")) \ .filter(F.col("date") == "2021-07-01") \ .groupBy("Grid_ID") \ .agg(F.collect_set("Latitude").getItem(0).cast("float").alias("corrected_lat"), F.collect_set("Longitude").getItem(0).cast("float").alias("corrected_long")) \ .withColumnRenamed("Grid_ID", "Grid_ID_dated") \ .select("Grid_ID_dated", "corrected_lat", "corrected_long") df_final = df.join(df_dated, on=[df.Grid_ID == df_dated.Grid_ID_dated], how="inner") \ .select(*df.columns, "corrected_lat", "corrected_long")
得到的输出符合预期:
Grid_ID Latitude Longitude DateTimeStamp corrected_lat corrected_long Grid_1 Lat1 Long1 2021-06-30 00:00:00 Lat2 Long2 Grid_1 Lat1 Long1 2021-06-30 00:01:00 Lat2 Long2 Grid_1 Lat1 Long1 2021-06-30 00:02:00 Lat2 Long2 Grid_1 Lat2 Long2 2021-07-01 00:00:00 Lat2 Long2 Grid_1 Lat2 Long2 2021-07-01 00:01:00 Lat2 Long2 Grid_1 Lat2 Long2 2021-07-01 00:02:00 Lat2 Long2
现在想请教:
- 是否可以用Window函数实现该需求?
- Window函数方案的效率是否优于当前的GroupBy+Join方式?
- 有没有更高效的实现方案?
解决方案
一、用Window函数实现
完全可以用Window函数实现,核心是基于Grid_ID分区,筛选出日期>2021-06-30的经纬度值,将其广播到同分区的所有行。代码如下:
import pyspark.sql.functions as F from pyspark.sql.window import Window # 定义Window:按Grid_ID分区 grid_window = Window.partitionBy("Grid_ID") df_final = df.withColumn("date", F.to_date("DateTimeStamp")) \ .withColumn("corrected_lat", F.first(F.when(F.col("date") > "2021-06-30", F.col("Latitude").cast("float")), ignorenulls=True).over(grid_window)) \ .withColumn("corrected_long", F.first(F.when(F.col("date") > "2021-06-30", F.col("Longitude").cast("float")), ignorenulls=True).over(grid_window)) \ .drop("date")
这里利用first(ignorenulls=True)自动取分区内符合条件的第一个非空值,正好对应每个Grid_ID的Lat2/Long2。
二、效率对比:Window vs GroupBy+Join
两者的效率差异取决于数据规模和集群配置:
- 小数据量:差异不明显,Window方案代码更简洁,无需额外Join操作,可能略快。
- 大数据量:
- GroupBy+Join需要先做Shuffle(GroupBy阶段),再做Join的Shuffle,两次Shuffle开销较高;
- Window函数的
first操作如果配合合适的分区策略(比如数据已按Grid_ID分区),可以避免额外Shuffle,仅做一次分区内计算,开销更低。
若Window函数触发全量Shuffle,两者开销接近。总体来说,Window方案在多数场景下比GroupBy+Join更高效,因为减少了一次Shuffle和Join的开销。
三、更高效的实现方案
1. 聚合后广播Join
每个Grid_ID对应的修正经纬度只有一组值,聚合后的小表数据量极小,可以将其广播,避免Join阶段的Shuffle:
import pyspark.sql.functions as F df_dated = df.withColumn("date", F.to_date("DateTimeStamp")) \ .filter(F.col("date") > "2021-06-30") \ .groupBy("Grid_ID") \ .agg(F.first("Latitude").cast("float").alias("corrected_lat"), F.first("Longitude").cast("float").alias("corrected_long")) # 广播小表,避免大表Shuffle df_final = df.join(F.broadcast(df_dated), on="Grid_ID", how="inner")
广播后Spark会将小表分发到各个Executor,无需对原大表做Shuffle,Join效率大幅提升,这是大数据场景下的最优选择之一。
2. 利用数据特性直接赋值
如果已知每个Grid_ID的Lat2/Long2对应日期是2021-07-01,且每个Grid_ID在该日期有数据,也可以用条件判断直接提取,但通用性稍弱:
import pyspark.sql.functions as F from pyspark.sql.window import Window grid_window = Window.partitionBy("Grid_ID") df_final = df.withColumn("corrected_lat", F.max(F.when(F.to_date("DateTimeStamp") == "2021-07-01", F.col("Latitude").cast("float"))).over(grid_window)) \ .withColumn("corrected_long", F.max(F.when(F.to_date("DateTimeStamp") == "2021-07-01", F.col("Longitude").cast("float"))).over(grid_window))
内容的提问来源于stack exchange,提问作者Skat
相关产品推荐
相关产品推荐

