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

如何用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

现在想请教:

  1. 是否可以用Window函数实现该需求?
  2. Window函数方案的效率是否优于当前的GroupBy+Join方式?
  3. 有没有更高效的实现方案?
解决方案

一、用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 15:19:21