Spark窗口函数rowsBetween仅保留完整窗口行的实现需求
I see you're calculating a moving median with a window spanning 50 rows before and after the current row, and you want to exclude rows where this window isn't fully populated (the first 50 and last 50 rows of each partition). Here's a straightforward way to adjust your approach:
Step 1: Track Row Position and Partition Size
First, we need to identify which rows have a complete window. We'll add two key columns to your dataset:
- A sequential row number ordered by
date_time_epochwithin each partition - The total number of rows in each partition
Step 2: Filter for Valid Rows
Once we have those values, we can filter out rows that don't have enough preceding or following rows. Specifically:
- Rows with a row number ≤ 50 (not enough rows before them)
- Rows with a row number > (total partition rows - 50) (not enough rows after them)
Step 3: Compute Moving Median on Filtered Data
Finally, apply your custom moving median UDAF to the filtered dataset to ensure only full windows are processed.
Full Modified Code
import org.apache.spark.sql.expressions.Window // Your existing MovingMedian UDAF (unchanged) class MovingMedian extends org.apache.spark.sql.expressions.UserDefinedAggregateFunction { def inputSchema: org.apache.spark.sql.types.StructType = org.apache.spark.sql.types.StructType( org.apache.spark.sql.types.StructField("value", org.apache.spark.sql.types.DoubleType) :: Nil ) def bufferSchema: org.apache.spark.sql.types.StructType = org.apache.spark.sql.types.StructType( org.apache.spark.sql.types.StructField("window_list", org.apache.spark.sql.types.ArrayType(org.apache.spark.sql.types.DoubleType, false)) :: Nil ) def dataType: org.apache.spark.sql.types.DataType = org.apache.spark.sql.types.DoubleType def deterministic: Boolean = true def initialize(buffer: org.apache.spark.sql.expressions.MutableAggregationBuffer): Unit = { buffer(0) = new scala.collection.mutable.ArrayBuffer[Double]() } def update(buffer: org.apache.spark.sql.expressions.MutableAggregationBuffer, input: org.apache.spark.sql.Row): Unit = { val bufferVal = buffer.getAs[scala.collection.mutable.WrappedArray[Double]](0).toBuffer bufferVal += input.getAs[Double](0) buffer(0) = bufferVal } def merge(buffer1: org.apache.spark.sql.expressions.MutableAggregationBuffer, buffer2: org.apache.spark.sql.Row): Unit = { buffer1(0) = buffer1.getAs[scala.collection.mutable.ArrayBuffer[Double]](0) ++ buffer2.getAs[scala.collection.mutable.ArrayBuffer[Double]](0) } def evaluate(buffer: org.apache.spark.sql.Row): Any = { val sortedWindow = buffer.getAs[scala.collection.mutable.WrappedArray[Double]](0).sorted.toBuffer val windowSize = sortedWindow.size if (windowSize % 2 == 0) { val index = windowSize / 2 (sortedWindow(index) + sortedWindow(index - 1)) / 2 } else { val index = (windowSize + 1) / 2 - 1 sortedWindow(index) } } } // Initialize the UDAF val mm = new MovingMedian // Define window specs for row number and total partition rows val partitionWindow = Window.partitionBy("raw_data_field_id") val orderedWindow = partitionWindow.orderBy("date_time_epoch") // Add row position and total rows metadata val dataWithRowInfo = rawdata .withColumn("row_num", org.apache.spark.sql.functions.row_number().over(orderedWindow)) .withColumn("total_rows", org.apache.spark.sql.functions.count("*").over(partitionWindow)) // Filter to keep only rows with full 50-row windows on both sides val filteredData = dataWithRowInfo .filter(org.apache.spark.sql.functions.col("row_num") > 50 && org.apache.spark.sql.functions.col("row_num") <= (org.apache.spark.sql.functions.col("total_rows") - 50)) // Calculate moving median on valid rows val rawdataFiltered = filteredData.withColumn("movingmedian", mm(col("value")).over( Window.partitionBy("raw_data_field_id").orderBy("date_time_epoch").rowsBetween(-50, 50)) )
Key Explanations
- Row Number: The
row_number()function assigns a unique position to each row in its partition, ordered by your timestamp. This lets us easily spot the early and late rows that can't have full windows. - Total Partition Rows: Using
count(*) over (partitionWindow)gives us the total size of each partition, so we know exactly how many rows at the end need to be excluded. - Filter Logic: The condition
row_num > 50androw_num <= total_rows - 50guarantees every remaining row has at least 50 rows before and after it, so yourrowsBetween(-50,50)window will always contain exactly 101 rows (current + 50 preceding + 50 following).
This approach ensures you only compute the moving median for valid, full-window rows, eliminating partial results at the start and end of each partition.
内容的提问来源于stack exchange,提问作者Remis Haroon - رامز

