Spark Scala分组取Top3并生成新列的实现求助
Spark Scala实现分组取TopN并转宽表
问题背景
现有如下DataFrame数据:
col1 col2 col3 requestErrorCode errorCodeCount 2023-02-27 LHR SFO 977 1 2023-02-27 LHR SFO 931 3 2023-02-27 ABC DEF 977 1 2023-02-27 ABC DEF 900 5 2023-02-27 ABC DEF 901 10 2023-02-27 ABC DEF 902 12 2023-02-27 ABC DEF 903 11 2023-02-27 ABC DEF 904 20 2023-02-27 GHI JKL 800 3 2023-02-27 GHI JKL 801 5 2023-02-27 GHI JKL 802 7 2023-02-27 GHI JKL 803 100 2023-02-27 GHI JKL 804 92 2023-02-27 GHI JKL 805 11 2023-02-27 GHI JKL 806 17
需要按col1、col2、col3分组,提取每组中errorCodeCount最高的前3条数据,转换为包含requestErrorCode1/errorCodeCount1等字段的宽表,预期结果如下:
col1 col2 col3 requestErrorCode1 errorCodeCount1 requestErrorCode2 errorCodeCount2 requestErrorCode3 errorCodeCount3 2023-02-27 LHR SFO 931 3 977 1 2023-02-27 ABC DEF 904 20 902 12 903 11 2023-02-27 GHI JKL 803 100 804 92 806 17
原有方法的问题
你之前尝试的orderBy+groupBy+first聚合只能获取每组的第一条数据,无法提取前3条,也无法将多条数据转换为多列结构,因此达不到预期效果。
正确实现步骤
1. 导入依赖包
import org.apache.spark.sql.expressions.Window import org.apache.spark.sql.functions.{desc, row_number, first}
2. 定义窗口函数生成组内排名
通过窗口函数给每个分组内的行按errorCodeCount降序分配排名:
// 按col1、col2、col3分组,按errorCodeCount降序排序 val windowSpec = Window.partitionBy("col1", "col2", "col3").orderBy(desc("errorCodeCount")) // 添加rank列,标记组内排名 val rankedDF = jointView.withColumn("rank", row_number().over(windowSpec))
3. 过滤出每组Top3数据
val top3DF = rankedDF.filter("rank <= 3")
4. 转换为宽表结构
通过分组聚合,将每个排名对应的字段提取并重命名为目标列:
val resultDF = top3DF.groupBy("col1", "col2", "col3") .agg( first("requestErrorCode").where("rank = 1").alias("requestErrorCode1"), first("errorCodeCount").where("rank = 1").alias("errorCodeCount1"), first("requestErrorCode").where("rank = 2").alias("requestErrorCode2"), first("errorCodeCount").where("rank = 2").alias("errorCodeCount2"), first("requestErrorCode").where("rank = 3").alias("requestErrorCode3"), first("errorCodeCount").where("rank = 3").alias("errorCodeCount3") )
5. 写入CSV(可选)
resultDF.write .option("header", "true") // 保留表头 .csv("/your/output/path")
灵活扩展(处理任意TopN)
如果需要调整取TopN的数量,可以动态生成聚合表达式,避免硬编码:
val topN = 3 // 可修改为任意数字 val aggExprs = (1 to topN).flatMap(rank => { Seq( first("requestErrorCode").where(s"rank = $rank").alias(s"requestErrorCode$rank"), first("errorCodeCount").where(s"rank = $rank").alias(s"errorCodeCount$rank") ) }) val resultDF = top3DF.groupBy("col1", "col2", "col3").agg(aggExprs.head, aggExprs.tail: _*)
说明
row_number()窗口函数确保每组内errorCodeCount最高的行排名为1,依次递减- 若分组内数据不足3条,对应的第N列会显示
null,与预期结果一致 - 最终生成的DataFrame结构完全匹配需求,可直接用于后续处理或写入文件
内容的提问来源于stack exchange,提问作者skd
相关产品推荐
相关产品推荐

