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

基于频率表的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 15:27:50