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

如何在Scala Spark中基于评分频率计算中位数?

在Spark/Scala中基于得票频率计算评分中位数

问题分析

你之前尝试用Array.fill($"votes_10")(10)生成评分数组失败,原因是**Array.fill是Scala本地集合的方法,仅支持编译时确定的整数常量,而col("votes_10")是Spark的Column对象(运行时才会解析的列值),两者不兼容**。

直接生成数组的方式在数据量大时还会导致内存占用过高,不适合Spark的分布式场景,更高效的方式是通过累计得票定位中位数区间。

推荐解决方案:累计得票法

核心思路:计算每个评分的累计得票数,找到第一个累计得票数超过中位数位置的评分,即为中位数(适合离散评分场景)。

完整代码示例

import org.apache.spark.sql.functions.{col, when}

val ratingsDfWithMedian = ratingsDf
  // 先计算总票数(如果还未生成)
  .withColumn("total_votes", 
    col("voted_1") + col("voted_2") + col("voted_3") + col("voted_4") + 
    col("voted_5") + col("voted_6") + col("voted_7") + col("voted_8") + 
    col("voted_9") + col("voted_10")
  )
  // 保留你原来的均值计算逻辑
  .withColumn("mean", 
    (col("voted_1")*1 + col("voted_2")*2 + col("voted_3")*3 + col("voted_4")*4 + 
     col("voted_5")*5 + col("voted_6")*6 + col("voted_7")*7 + col("voted_8")*8 + 
     col("voted_9")*9 + col("voted_10")*10) / col("total_votes")
  )
  // 计算中位数位置(奇数取中间位,偶数取中间两个的前一位,符合离散评分的中位数定义)
  .withColumn("median_pos", (col("total_votes") + 1) / 2.0)
  // 从高到低计算累计得票数
  .withColumn("cum_10", col("voted_10"))
  .withColumn("cum_9", col("cum_10") + col("voted_9"))
  .withColumn("cum_8", col("cum_9") + col("voted_8"))
  .withColumn("cum_7", col("cum_8") + col("voted_7"))
  .withColumn("cum_6", col("cum_7") + col("voted_6"))
  .withColumn("cum_5", col("cum_6") + col("voted_5"))
  .withColumn("cum_4", col("cum_5") + col("voted_4"))
  .withColumn("cum_3", col("cum_4") + col("voted_3"))
  .withColumn("cum_2", col("cum_3") + col("voted_2"))
  .withColumn("cum_1", col("cum_2") + col("voted_1"))
  // 判断中位数所在的评分区间
  .withColumn("median", 
    when(col("cum_10") >= col("median_pos"), 10)
    .when(col("cum_9") >= col("median_pos"), 9)
    .when(col("cum_8") >= col("median_pos"), 8)
    .when(col("cum_7") >= col("median_pos"), 7)
    .when(col("cum_6") >= col("median_pos"), 6)
    .when(col("cum_5") >= col("median_pos"), 5)
    .when(col("cum_4") >= col("median_pos"), 4)
    .when(col("cum_3") >= col("median_pos"), 3)
    .when(col("cum_2") >= col("median_pos"), 2)
    .otherwise(1)
  )
  // 清理中间临时列
  .drop("median_pos", "cum_1", "cum_2", "cum_3", "cum_4", "cum_5", "cum_6", "cum_7", "cum_8", "cum_9", "cum_10")

逻辑说明

  1. 总票数与均值:先计算所有评分的总得票数,再复用你原来的均值计算逻辑。
  2. 中位数位置:使用(total_votes + 1)/2.0计算,确保奇数和偶数总票数都能定位到对应的中位数区间。
  3. 累计得票:从最高评分(10)开始累加得票数,直到累计值超过中位数位置,当前评分即为中位数。
  4. 清理临时列:最后去掉中间生成的累计列和位置列,保持数据表简洁。

备选方案:UDF生成评分数组(不推荐大数据场景)

如果一定要用数组方式计算(仅适合小数据量),可以自定义UDF生成评分数组后计算中位数:

import org.apache.spark.sql.functions.udf

// 自定义UDF,接收各评分得票数,生成评分数组并计算中位数
def getMedian(v1: Int, v2: Int, v3: Int, v4: Int, v5: Int, v6: Int, v7: Int, v8: Int, v9: Int, v10: Int): Int = {
  val ratingsList = List.fill(v10)(10) ++ List.fill(v9)(9) ++ List.fill(v8)(8) ++ 
                    List.fill(v7)(7) ++ List.fill(v6)(6) ++ List.fill(v5)(5) ++ 
                    List.fill(v4)(4) ++ List.fill(v3)(3) ++ List.fill(v2)(2) ++ List.fill(v1)(1)
  val len = ratingsList.length
  if (len % 2 == 1) ratingsList(len / 2) else (ratingsList(len/2 -1) + ratingsList(len/2)) / 2
}

val medianUdf = udf(getMedian _)

val ratingsDfWithMedian = ratingsDf
  .withColumn("median", medianUdf(
    col("voted_1"), col("voted_2"), col("voted_3"), col("voted_4"), 
    col("voted_5"), col("voted_6"), col("voted_7"), col("voted_8"), 
    col("voted_9"), col("voted_10")
  ))

⚠️ 注意:这种方法会生成包含所有评分的大列表,数据量大时极易导致内存溢出,不适合Spark分布式处理的场景。


内容的提问来源于stack exchange,提问作者Plk

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 00:02:55