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

如何在Spark中结合Window rowsBetween实现自定义聚合器?

基于Spark窗口函数的自定义速度计算实现

优先方案:用lag函数替代自定义聚合器

你的窗口范围是rowsBetween(-1, currentRow),本质是取当前行和上一行数据,用Spark内置的lag函数直接获取上一行的时间、经纬度,再通过UDF计算速度,比自定义聚合器更简洁高效,代码如下:

import org.apache.spark.sql.functions._
import org.apache.spark.sql.types._

// 实现速度计算UDF:输入当前行和上一行的时间、经纬度,返回米/秒单位的速度
val calculateSpeed = udf((currTs: Long, currLat: Double, currLon: Double, prevTs: Long, prevLat: Double, prevLon: Double) => {
  if (prevTs == 0 || currTs == prevTs) 0.0
  else {
    // Haversine公式计算两点间距离(单位:米)
    val earthRadius = 6371000
    val dLat = Math.toRadians(currLat - prevLat)
    val dLon = Math.toRadians(currLon - prevLon)
    val a = Math.sin(dLat/2) * Math.sin(dLat/2) +
            Math.cos(Math.toRadians(prevLat)) * Math.cos(Math.toRadians(currLat)) *
            Math.sin(dLon/2) * Math.sin(dLon/2)
    val c = 2 * Math.atan2(Math.sqrt(a), Math.sqrt(1-a))
    val distance = earthRadius * c
    // 计算时间差(秒,假设timestamp是毫秒级)
    val timeDiff = (currTs - prevTs) / 1000.0
    distance / timeDiff
  }
})

// 定义窗口规则
val trackWindow = Window
  .partitionBy(trackIdColumn)
  .orderBy("timestamp")

// 生成速度列
retDf = retDf
  .withColumn("prev_timestamp", lag("timestamp", 1).over(trackWindow))
  .withColumn("prev_latitude", lag("latitude", 1).over(trackWindow))
  .withColumn("prev_longitude", lag("longitude", 1).over(trackWindow))
  .withColumn("speed", calculateSpeed(col("timestamp"), col("latitude"), col("longitude"), col("prev_timestamp"), col("prev_latitude"), col("prev_longitude")))
  .drop("prev_timestamp", "prev_latitude", "prev_longitude")

自定义窗口聚合器实现(按需使用)

如果必须用自定义聚合器,可基于Aggregator类实现(DeclarativeAggregate依赖表达式体系,实现成本更高),具体步骤如下:

1. 定义数据样例类

// 单条GPS数据的样例类
case class GpsPoint(timestamp: Long, latitude: Double, longitude: Double)
// 聚合过程的状态样例类,用于保存窗口内的两行数据
case class SpeedAggState(point1: Option[GpsPoint], point2: Option[GpsPoint])

2. 实现自定义Aggregator

import org.apache.spark.sql.expressions.Aggregator
import org.apache.spark.sql.{Encoder, Encoders}

object SpeedAggregator extends Aggregator[GpsPoint, SpeedAggState, Double] {
  // 初始化状态:无数据时的默认状态
  override def zero: SpeedAggState = SpeedAggState(None, None)

  // 更新状态:窗口内每加入一行数据,更新状态存储的点
  override def reduce(state: SpeedAggState, input: GpsPoint): SpeedAggState = {
    state.point1 match {
      case None => SpeedAggState(Some(input), None)
      case Some(_) => SpeedAggState(state.point1, Some(input))
    }
  }

  // 合并状态:窗口聚合场景下基本用不到,按规则实现即可
  override def merge(state1: SpeedAggState, state2: SpeedAggState): SpeedAggState = {
    (state1.point1, state1.point2) match {
      case (None, None) => state2
      case (p1, None) => SpeedAggState(p1, state2.point1)
      case _ => state1
    }
  }

  // 计算最终速度:当状态内有两行数据时,用Haversine公式计算速度
  override def finish(reduction: SpeedAggState): Double = {
    (reduction.point1, reduction.point2) match {
      case (Some(prev), Some(curr)) =>
        val earthRadius = 6371000
        val dLat = Math.toRadians(curr.latitude - prev.latitude)
        val dLon = Math.toRadians(curr.longitude - prev.longitude)
        val a = Math.sin(dLat/2) * Math.sin(dLat/2) +
                Math.cos(Math.toRadians(prev.latitude)) * Math.cos(Math.toRadians(curr.latitude)) *
                Math.sin(dLon/2) * Math.sin(dLon/2)
        val c = 2 * Math.atan2(Math.sqrt(a), Math.sqrt(1-a))
        val distance = earthRadius * c
        val timeDiff = (curr.timestamp - prev.timestamp) / 1000.0
        if (timeDiff <= 0) 0.0 else distance / timeDiff
      case _ => 0.0
    }
  }

  // 状态和输出结果的编码器
  override def bufferEncoder: Encoder[SpeedAggState] = Encoders.product[SpeedAggState]
  override def outputEncoder: Encoder[Double] = Encoders.scalaDouble
}

3. 调用自定义聚合器

import org.apache.spark.sql.functions._

// 将三列打包为GpsPoint类型
val gpsPointCol = struct(
  col("timestamp").cast(LongType),
  col("latitude").cast(DoubleType),
  col("longitude").cast(DoubleType)
).as("gps_point")

// 定义窗口规则
val trackWindow = Window
  .partitionBy(trackIdColumn)
  .orderBy("timestamp")
  .rowsBetween(-1, Window.currentRow)

// 生成速度列
retDf = retDf
  .withColumn("gps_point", gpsPointCol)
  .withColumn("speed", callUDF(SpeedAggregator, col("gps_point")).over(trackWindow))
  .drop("gps_point")

异常值剔除方案

计算出速度后,可通过以下方式剔除异常值:

  • 固定阈值过滤:设定合理的速度上限(比如120km/h=33.3m/s),过滤超阈值数据
    val filteredDf = retDf.filter(col("speed") <= 33.3 || col("speed").isNull)
    
  • 统计法过滤:基于轨迹内的速度均值和标准差,剔除偏离3倍标准差的数据
    val statsWindow = Window.partitionBy(trackIdColumn)
    val filteredDf = retDf
      .withColumn("speed_mean", avg(col("speed")).over(statsWindow))
      .withColumn("speed_stddev", stddev(col("speed")).over(statsWindow))
      .filter(col("speed").between(col("speed_mean") - 3*col("speed_stddev"), col("speed_mean") + 3*col("speed_stddev")))
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 15:32:04