Spark Scala自定义UserDefinedAggregateFunction的merge缓冲区为空问题
Hey, let's break down what's happening here — your merge function is only seeing initialized buffer values because your update method isn't actually being called at all! Here's why and how to fix it:
核心问题:输入列类型不匹配
When you load the CSV file without specifying a schema or casting columns, Spark reads all columns as StringType by default. But your UDAF's inputSchema expects two DoubleType values (min and max). Spark silently skips calling the update method when the input types don't match the UDAF's expected schema, so your buffer never gets populated with real data. That's why you only see the initial 0/0.0 values in merge.
解决方案步骤
1. 转换输入列的类型为Double
Modify your DataFrame loading code to cast MinPrice and MaxPrice to DoubleType before using the UDAF:
val df = spark.read.format("csv") .option("header", "true") .load("./data/2018-02-01/2018-02-01_BINS_XETR08.csv") .withColumn("MinPrice", col("MinPrice").cast(DoubleType)) .withColumn("MaxPrice", col("MaxPrice").cast(DoubleType))
Alternatively, you can define a schema upfront to avoid casting:
val csvSchema = StructType(Array( // Add all your CSV columns here, including MinPrice and MaxPrice as DoubleType StructField("Mnemonic", StringType), StructField("Date", StringType), StructField("MinPrice", DoubleType), StructField("MaxPrice", DoubleType), // ... other columns specific to your CSV )) val df = spark.read.format("csv") .option("header", "true") .schema(csvSchema) .load("./data/2018-02-01/2018-02-01_BINS_XETR08.csv")
2. 简化UDF注册(可选但推荐)
You don't need to register the UDAF with both sqlContext and spark — using the SparkSession's udf.register is sufficient for modern Spark versions:
val adv = new AggregateDeltaVolume spark.udf.register("adv", adv)
3. 添加调试打印验证update是否执行
Add a debug print in your update method to confirm it's being called after fixing the type issue:
override def update(buffer: MutableAggregationBuffer, input: Row): Unit = { println("DEBUG: AggregateDeltaVolume.update called with new data!") val min:Double = input.getAs[Double](0) val max:Double = input.getAs[Double](1) val delta:Double = (max - min) / min buffer(0) = buffer.getAs[Long](0) + 1L buffer(1) = buffer.getAs[Double](1) + delta }
You should now see this print statement in your logs, confirming real data is being processed.
4. 优化merge方法的调试(可选)
Update your merge debug print to include the input buffer values too — this will help you verify that data is being passed between nodes:
override def merge(buffer: MutableAggregationBuffer, input: Row): Unit = { println(s"DEBUG: AggregateDeltaVolume.merge - target buffer: (${buffer(0)}, ${buffer(1)}), input buffer: (${input(0)}, ${input(1)})") buffer(0) = buffer.getAs[Long](0) + input.getAs[Long](0) buffer(1) = buffer.getAs[Double](1) + input.getAs[Double](1) }
最终修改后的核心代码片段
// UDAF class remains the same class AggregateDeltaVolume extends UserDefinedAggregateFunction { override def inputSchema: StructType = StructType( Array(StructField("min", DoubleType), StructField("max", DoubleType)) ) override def bufferSchema: StructType = StructType( StructField("count", LongType) :: StructField("volumeDeltaSum", DoubleType) :: Nil ) override def dataType: DataType = DoubleType override def initialize(buffer: MutableAggregationBuffer): Unit = { buffer(0) = 0L buffer(1) = 0.0d } override def update(buffer: MutableAggregationBuffer, input: Row): Unit = { println("DEBUG: AggregateDeltaVolume.update called with new data!") val min:Double = input.getAs[Double](0) val max:Double = input.getAs[Double](1) val delta:Double = (max - min) / min buffer(0) = buffer.getAs[Long](0) + 1L buffer(1) = buffer.getAs[Double](1) + delta } override def merge(buffer: MutableAggregationBuffer, input: Row): Unit = { println(s"DEBUG: AggregateDeltaVolume.merge - target buffer: (${buffer(0)}, ${buffer(1)}), input buffer: (${input(0)}, ${input(1)})") buffer(0) = buffer.getAs[Long](0) + input.getAs[Long](0) buffer(1) = buffer.getAs[Double](1) + input.getAs[Double](1) } override def evaluate(buffer: Row): Any = { buffer.getDouble(1) / buffer.getLong(0) } override def deterministic: Boolean = true } // Initialize SparkSession (modern alternative to sqlContext) val spark = SparkSession.builder() .appName("StockAnalysis") .getOrCreate() import spark.implicits._ import org.apache.spark.sql.functions._ val adv = new AggregateDeltaVolume spark.udf.register("adv", adv) // Load CSV with type casting val df = spark.read.format("csv") .option("header", "true") .load("./data/2018-02-01/2018-02-01_BINS_XETR08.csv") .withColumn("MinPrice", col("MinPrice").cast(DoubleType)) .withColumn("MaxPrice", col("MaxPrice").cast(DoubleType)) val dayView = df.groupBy("Mnemonic", "Date") .agg(expr("adv(MinPrice, MaxPrice) as AvgDeltaVolume")) dayView.show()
After making these changes, your UDAF should correctly populate the buffer in update, and the merge function will receive non-default values from other partitions.
内容的提问来源于stack exchange,提问作者user1792160

