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

在Scala中实现分组条件化80%截尾均值的方法

解决Scala Spark中分组条件计算截尾均值/均值的问题

好的,我来帮你搞定这个需求!首先得纠正一下你当前代码里的小问题:groupBy("city", "year").count()之后再去agg(avg($"amount"))是行不通的,因为count()会把原表的amount列给丢掉了,咱们得换个思路来实现。

先明确核心需求:按city和year分组后,如果分组记录数≥10,就计算80%截尾均值(也就是丢弃首尾各10%的数据,取中间80%的均值);否则直接用普通均值。

先理清楚80%截尾均值的计算逻辑

按照你说的规则,80%截尾均值的步骤是:

  • 把分组内的amount按从小到大排序
  • 算出要丢弃的首尾数据量:分组总条数的10%(因为要留中间80%),这里注意取整的问题,比如分组有15条数据,10%是1.5,咱们可以用floor取1,也就是首尾各丢1条,留中间13条
  • 去掉排序后的前N条和后N条,剩下的部分算均值就行

具体实现步骤

咱们可以用窗口函数先给每条数据标记分组信息,再根据分组大小选择计算方式,下面是完整代码:

完整可运行代码

import org.apache.spark.sql.expressions.Window
import org.apache.spark.sql.functions._

// 你的示例数据
val sales = Seq( 
  ("Warsaw", 2016, 100), ("Warsaw", 2017, 200), 
  ("Boston", 2015, 50), ("Boston", 2016, 150), 
  ("Toronto", 2017, 50)
).toDF("city", "year", "amount")

// 定义两个窗口:一个用来给分组内的amount排序,一个用来计算分组总条数
val sortWindow = Window.partitionBy("city", "year").orderBy("amount")
val groupSizeWindow = Window.partitionBy("city", "year")

// 给每条数据加上分组总条数和排序后的行号
val withMeta = sales
  .withColumn("group_size", count("*").over(groupSizeWindow))
  .withColumn("row_num", row_number().over(sortWindow))

// 计算截尾均值:筛选出中间80%的数据后再算均值
val trimmedAvgDF = withMeta
  .withColumn("trim_ratio", lit(0.1)) // 首尾各去掉10%,对应保留80%
  .withColumn("start_row", floor($"group_size" * $"trim_ratio") + 1) // 起始行(包含)
  .withColumn("end_row", $"group_size" - floor($"group_size" * $"trim_ratio")) // 结束行(包含)
  .filter($"row_num".between($"start_row", $"end_row"))
  .groupBy("city", "year", "group_size")
  .agg(avg($"amount").as("trimmed_avg"))

// 计算所有分组的普通均值
val regularAvgDF = sales
  .groupBy("city", "year")
  .agg(count("*").as("group_size"), avg($"amount").as("regular_avg"))

// 合并结果:根据分组大小选择用截尾均值还是普通均值
val finalResult = regularAvgDF
  .join(trimmedAvgDF, Seq("city", "year", "group_size"), "left")
  .withColumn("final_avg", when($"group_size" >= 10, $"trimmed_avg").otherwise($"regular_avg"))
  .select("city", "year", "group_size", "final_avg")

// 查看结果
finalResult.show()

代码细节解释

  1. 窗口函数的作用:
    • sortWindow:给每个city+year分组内的amount排序,生成行号,方便后续筛选中间数据
    • groupSizeWindow:专门用来计算每个分组的总记录数,这样我们就能判断是否需要用截尾均值
  2. 截尾均值的筛选逻辑:
    • 我这里用trim_ratio = 0.1,对应首尾各去10%,如果你需要调整截尾比例,改这个值就行
    • 用floor处理取整,如果你希望更严格的截尾(比如分组11条时去掉首尾各2条),可以换成ceil
  3. 结果合并:
    • 先算所有分组的普通均值,再和截尾均值的结果左连接(因为只有分组≥10的才有截尾均值)
    • 用when函数做条件判断,自动选择对应的均值

更贴合百分比的替代方案

如果分组数据量很大,你可能想要更精确的按百分比截尾,而不是按行号,这时候可以用percent_rank()窗口函数替代row_number():

// 用percent_rank实现的截尾均值计算
val withMetaPct = sales
  .withColumn("group_size", count("*").over(groupSizeWindow))
  .withColumn("pct_rank", percent_rank().over(sortWindow))

val trimmedAvgDFPct = withMetaPct
  .filter($"pct_rank".between(0.1, 0.9)) // 直接保留10%到90%分位之间的数据
  .groupBy("city", "year", "group_size")
  .agg(avg($"amount").as("trimmed_avg"))

这种方式更精准,尤其是当分组数据量很大的时候,不会因为取整问题影响截尾比例。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.27 07:34:14