在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()
代码细节解释
- 窗口函数的作用:
sortWindow:给每个city+year分组内的amount排序,生成行号,方便后续筛选中间数据groupSizeWindow:专门用来计算每个分组的总记录数,这样我们就能判断是否需要用截尾均值
- 截尾均值的筛选逻辑:
- 我这里用
trim_ratio = 0.1,对应首尾各去10%,如果你需要调整截尾比例,改这个值就行 - 用
floor处理取整,如果你希望更严格的截尾(比如分组11条时去掉首尾各2条),可以换成ceil
- 我这里用
- 结果合并:
- 先算所有分组的普通均值,再和截尾均值的结果左连接(因为只有分组≥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
相关产品推荐
相关产品推荐

