基于频率表的Spark UDAF实现加权统计量计算
实现Spark加权统计自定义UDAF(针对带频率的分组统计需求)
针对候选人测试得分带频率的场景,由于Spark内置的approxQuantile不支持加权计算,展开数据的方式效率低下,我们可以通过自定义UDAF(用户定义聚合函数)来高效完成所有要求的加权统计量计算。
一、核心思路
自定义UDAF在聚合阶段直接处理加权数据,无需展开重复记录,通过维护以下中间状态完成计算:
- 总频率、加权和、加权平方和:用于计算加权均值与标准差
- 排序后的(得分、累计频率)对:用于计算分位数、中位数
- 众数对应的得分与最高频率:直接跟踪众数
二、自定义UDAF实现
1. 缓冲类定义
用于存储聚合过程中的中间数据:
import org.apache.spark.sql.expressions.{MutableAggregationBuffer, UserDefinedAggregateFunction} import org.apache.spark.sql.types._ import org.apache.spark.sql.Row case class WeightedStatsBuffer( var totalFreq: Long, var weightedSum: Double, var weightedSumSq: Double, var scoreFreqPairs: List[(Double, Long)], var modeScore: Double, var modeFreq: Long )
2. UDAF实现类
继承UserDefinedAggregateFunction完成完整聚合逻辑:
class WeightedStatsUDAF extends UserDefinedAggregateFunction { // 输入结构:(得分: Double, 频率: Long) override def inputSchema: StructType = StructType( StructField("score", DoubleType) :: StructField("frequency", LongType) :: Nil ) // 缓冲数据结构 override def bufferSchema: StructType = StructType( StructField("totalFreq", LongType) :: StructField("weightedSum", DoubleType) :: StructField("weightedSumSq", DoubleType) :: StructField("scoreFreqPairs", ArrayType(StructType( StructField("score", DoubleType) :: StructField("freq", LongType) :: Nil ))) :: StructField("modeScore", DoubleType) :: StructField("modeFreq", LongType) :: Nil ) // 输出结果结构:包含所有要求的统计量 override def dataType: DataType = StructType( StructField("weightedMean", DoubleType) :: StructField("weightedMedian", DoubleType) :: StructField("weightedMode", DoubleType) :: StructField("weightedStdDev", DoubleType) :: StructField("p10", DoubleType) :: StructField("p90", DoubleType) :: StructField("trimmedMean", DoubleType) :: StructField("trimmedStdDev", DoubleType) :: Nil ) override def deterministic: Boolean = true // 初始化缓冲 override def initialize(buffer: MutableAggregationBuffer): Unit = { buffer(0) = 0L buffer(1) = 0.0 buffer(2) = 0.0 buffer(3) = Array.empty[(Double, Long)] buffer(4) = 0.0 buffer(5) = 0L } // 更新单条输入数据到缓冲 override def update(buffer: MutableAggregationBuffer, input: Row): Unit = { val score = input.getDouble(0) val freq = input.getLong(1) if (freq > 0) { // 更新基础统计量 buffer(0) = buffer.getLong(0) + freq buffer(1) = buffer.getDouble(1) + score * freq buffer(2) = buffer.getDouble(2) + score * score * freq // 追加得分-频率对 val currentPairs = buffer.getAs[Array[(Double, Long)]](3).toList buffer(3) = (currentPairs :+ (score, freq)).toArray // 更新众数 val currentModeFreq = buffer.getLong(5) if (freq > currentModeFreq) { buffer(4) = score buffer(5) = freq } else if (freq == currentModeFreq && score < buffer.getDouble(4)) { buffer(4) = score } } } // 合并两个缓冲数据 override def merge(buffer1: MutableAggregationBuffer, buffer2: Row): Unit = { // 合并基础统计量 buffer1(0) = buffer1.getLong(0) + buffer2.getLong(0) buffer1(1) = buffer1.getDouble(1) + buffer2.getDouble(1) buffer1(2) = buffer1.getDouble(2) + buffer2.getDouble(2) // 合并得分-频率对 val pairs1 = buffer1.getAs[Array[(Double, Long)]](3).toList val pairs2 = buffer2.getAs[Array[(Double, Long)]](3).toList buffer1(3) = (pairs1 ++ pairs2).toArray // 合并众数 val modeFreq1 = buffer1.getLong(5) val modeFreq2 = buffer2.getLong(5) if (modeFreq2 > modeFreq1) { buffer1(4) = buffer2.getDouble(4) buffer1(5) = modeFreq2 } else if (modeFreq2 == modeFreq1 && buffer2.getDouble(4) < buffer1.getDouble(4)) { buffer1(4) = buffer2.getDouble(4) } } // 计算最终统计结果 override def evaluate(buffer: Row): Any = { val totalFreq = buffer.getLong(0) if (totalFreq == 0) { Row(null, null, null, null, null, null, null, null) } else { // 加权均值与标准差 val weightedMean = buffer.getDouble(1) / totalFreq val weightedVar = (buffer.getDouble(2) / totalFreq) - math.pow(weightedMean, 2) val weightedStdDev = math.sqrt(weightedVar) // 处理得分-频率对:去重、排序、计算累计频率 val sortedPairs = buffer.getAs[Array[(Double, Long)]](3) .groupBy(_._1) .mapValues(_.map(_._2).sum) .toList .sortBy(_._1) val cumulativeFreqs = sortedPairs.scanLeft(0L)((acc, pair) => acc + pair._2).tail // 计算中位数、10/90分位数 val median = findQuantile(sortedPairs, cumulativeFreqs, totalFreq * 0.5) val p10 = findQuantile(sortedPairs, cumulativeFreqs, totalFreq * 0.1) val p90 = findQuantile(sortedPairs, cumulativeFreqs, totalFreq * 0.9) // 计算10-90分位内的加权统计量 val (trimmedSum, trimmedSumSq, trimmedFreq) = sortedPairs.zip(cumulativeFreqs).foldLeft((0.0, 0.0, 0L)) { case ((sum, sumSq, freq), ((score, cnt), cumFreq)) => val prevCumFreq = if (cumFreq - cnt == 0) 0 else cumulativeFreqs(cumulativeFreqs.indexOf(cumFreq) - 1) if (cumFreq <= p10 || prevCumFreq >= p90) { (sum, sumSq, freq) } else if (prevCumFreq <= p10 && cumFreq >= p90) { val takeStart = (p10 - prevCumFreq).toLong val takeEnd = (p90 - prevCumFreq).toLong (sum + score * (takeEnd - takeStart), sumSq + score * score * (takeEnd - takeStart), freq + (takeEnd - takeStart)) } else if (prevCumFreq <= p10) { val takeFreq = (cumFreq - p10).toLong (sum + score * takeFreq, sumSq + score * score * takeFreq, freq + takeFreq) } else if (cumFreq >= p90) { val takeFreq = (p90 - prevCumFreq).toLong (sum + score * takeFreq, sumSq + score * score * takeFreq, freq + takeFreq) } else { (sum + score * cnt, sumSq + score * score * cnt, freq + cnt) } } val trimmedMean = if (trimmedFreq > 0) trimmedSum / trimmedFreq else null val trimmedVar = if (trimmedFreq > 0) (trimmedSumSq / trimmedFreq) - math.pow(trimmedMean, 2) else null val trimmedStdDev = if (trimmedVar != null) math.sqrt(trimmedVar) else null // 组装结果 Row(weightedMean, median, buffer.getDouble(4), weightedStdDev, p10, p90, trimmedMean, trimmedStdDev) } } // 辅助方法:计算指定位置的分位数 private def findQuantile(sortedPairs: List[(Double, Long)], cumulativeFreqs: List[Long], pos: Double): Double = { sortedPairs.zip(cumulativeFreqs).find { case ((_, cnt), cumFreq) => cumFreq >= pos } match { case Some(((score, _), cumFreq)) => val prevCumFreq = if (cumFreq - cnt == 0) 0 else cumulativeFreqs(cumulativeFreqs.indexOf(cumFreq) - 1) // 线性插值优化精度,可根据需求调整为直接取当前得分 val lowerScore = sortedPairs(cumulativeFreqs.indexOf(prevCumFreq))._1 lowerScore + (score - lowerScore) * (pos - prevCumFreq) / (cumFreq - prevCumFreq) case None => 0.0 } } }
三、使用UDAF进行分组统计
注册UDAF后,直接应用到分组数据上:
// 注册UDAF spark.udf.register("weightedStats", new WeightedStatsUDAF()) // 按候选人、测试ID分组计算 val resultDF = df.groupBy("candidate", "testID") .agg(callUDF("weightedStats", col("score"), col("Frequency")).alias("stats")) .select( col("candidate"), col("testID"), col("stats.weightedMean").alias("weighted_mean"), col("stats.weightedMedian").alias("weighted_median"), col("stats.weightedMode").alias("weighted_mode"), col("stats.weightedStdDev").alias("weighted_std_dev"), col("stats.p10").alias("p10"), col("stats.p90").alias("p90"), col("stats.trimmedMean").alias("trimmed_mean"), col("stats.trimmedStdDev").alias("trimmed_std_dev") ) resultDF.show()
四、优化与调整说明
- 分位数计算:当前实现使用线性插值提升精度,若无需高精度可直接取第一个累计频率超过目标位置的得分。
- 众数逻辑:当前在频率相同时取较小得分,可根据业务需求修改为其他规则(如取较大得分)。
- 空值处理:分组内无有效数据时返回null,可根据需求替换为默认值(如0.0)。
内容的提问来源于stack exchange,提问作者user2458922
相关产品推荐
相关产品推荐

