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

Spark窗口函数rowsBetween仅保留完整窗口行的实现需求

Solution to Exclude Incomplete Windows for Spark Moving Median

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_epoch within 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 > 50 and row_num <= total_rows - 50 guarantees every remaining row has at least 50 rows before and after it, so your rowsBetween(-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 - رامز

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.13 07:40:02