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

