Scala中Spark DataFrame按ID分组按规则取多条记录的实现问题
实现方案
核心思路是用Spark窗口函数分别处理两类筛选规则,合并后得到最终结果,完整代码如下:
import org.apache.spark.sql.expressions.Window import org.apache.spark.sql.functions._ // 1. 给每个id分组标记分组内的最小hierarchy val wId = Window.partitionBy("id") val dfWithMinHier = df.withColumn("min_hierarchy", min("hierarchy").over(wId)) // 2. 处理最低层级的筛选规则 val wMinHier = Window.partitionBy("id", "hierarchy").orderBy(col("amount").desc) val lowHierPart = dfWithMinHier .filter(col("hierarchy") === col("min_hierarchy")) .withColumn("rn", row_number().over(wMinHier)) .filter( (col("min_hierarchy") === 1 && col("rn") <= 1) || (col("min_hierarchy") === 2 && col("rn") <= 2) ) .drop("min_hierarchy", "rn") // 3. 处理hierarchy>=4的筛选规则 val wHighHier = Window.partitionBy("id").orderBy(col("amount").desc) val highHierPart = dfWithMinHier .filter(col("hierarchy") >= 4) .withColumn("rn", row_number().over(wHighHier)) .filter(col("rn") <= 3) .drop("min_hierarchy", "rn") // 4. 合并两部分结果,去重避免边界重叠 val result = lowHierPart.unionByName(highHierPart).dropDuplicates() // 查看结果,可按需排序 result.orderBy("id", "hierarchy", col("amount").desc).show()
代码说明
- 所有计算复用了带min_hierarchy标记的中间表,避免多次扫描原始数据产生额外shuffle
- 用row_number窗口函数实现按金额倒序取TopN的逻辑,符合筛选要求
- union后做去重处理,覆盖min_hierarchy>=4的极端边界场景,避免重复记录
- 如果需要匹配示例输出中hierarchy>=4的每个层级单独取前3的效果,只需将
wHighHier的分区规则改为Window.partitionBy("id", "hierarchy")即可
内容的提问来源于stack exchange,提问作者Rolando
相关产品推荐
相关产品推荐

