Spark 3.0.1(Scala):如何在Shuffle前对分组行执行自定义操作?
Spark 3.0.1 + Scala 自定义UDAF实现分组筛选与TopN问题
我用Spark 3.0.1和Scala处理一个包含大量相似行的数据集,需要完成三个操作:
- 在各分区内按键对行分组
- 从每个分区内的每组中选择行的子集
- 跨分区获取每组中按某列值排序的顶部行
用Spark内置API实现这类操作比较简单,示例代码如下:
val df2 = df1 .groupBy("group_key", df1.columns: _*) .min("length") df2 .show()
但我搞不定自定义UDAF的语法,自己写的代码一直跑不起来,代码如下:
def similar(l1: List[Row], l2: List[Row]): List[Row] = { // 这里是筛选非相似行并限制列表大小以避免分区失衡的逻辑 l1 ++ l2 } val aggFuntion = new Aggregator[Row, List[Row], List[Row]]( createCombiner = List[Row], mergeValue = similar, mergeCombiners = similar ) val aggUdaf = udaf[Row, List[Row], List[Row]](aggFuntion) val df2 = df1 .groupBy("group_key", df1.columns: _*) .agg(aggUdaf(col("fathers"))) df2 .show()
原代码的问题
- Aggregator实例化方式错了:Scala里
Aggregator是抽象类,不能直接new并传参数,得继承它实现对应方法。 - createCombiner定义错误:它需要是把单个
Row转成中间状态List[Row]的函数,不是直接写List[Row],应该写成(row: Row) => List(row)。 - mergeValue方法签名不对:这个方法是把单个输入行合并到缓冲区,参数应该是
(buffer: List[Row], input: Row),不是两个List[Row]。 - 缺少Encoder定义:
Aggregator必须为中间状态和输出类型提供Encoder,要么重写对应方法,要么用隐式上下文的Encoder。 - groupBy逻辑不合理:按
group_key加所有列分组,等于每行单独成组,完全不符合分组筛选的需求,应该只按group_key分组。
修正后的UDAF实现
下面是能正常运行的修正版代码,实现分区内筛选、跨分区取TopN的逻辑:
import org.apache.spark.sql.{Encoder, Encoders} import org.apache.spark.sql.expressions.Aggregator import org.apache.spark.sql.Row // 自定义Aggregator,实现分组内收集行并筛选,最终输出筛选后的行列表 class SimilarRowAggregator(limit: Int) extends Aggregator[Row, List[Row], List[Row]] { // 创建初始缓冲区:把单个Row转成List override def createCombiner(row: Row): List[Row] = List(row) // 将单个输入Row合并到缓冲区:替换成你的相似行筛选逻辑即可 override def mergeValue(buffer: List[Row], row: Row): List[Row] = { // 示例逻辑:行不在缓冲区且列表未达上限时才添加,可替换成你的相似性判断 if (!buffer.contains(row) && buffer.size < limit) buffer :+ row else buffer } // 合并两个缓冲区:去重并限制大小,避免数据膨胀 override def mergeCombiners(buffer1: List[Row], buffer2: List[Row]): List[Row] = { val combined = buffer1 ++ buffer2 // 去重后取前limit行,也可以先按目标列排序再截取 combined.distinct.take(limit) } // 最终输出筛选后的列表 override def finish(buffer: List[Row]): List[Row] = buffer // 中间状态的Encoder:用kryo序列化Row的List override def bufferEncoder: Encoder[List[Row]] = Encoders.kryo[List[Row]] // 输出类型的Encoder:同上 override def outputEncoder: Encoder[List[Row]] = Encoders.kryo[List[Row]] } // 将Aggregator转为可在DataFrame中使用的列函数 val similarRowUdaf = new SimilarRowAggregator(limit = 10).toColumn // 正确使用UDAF做分组聚合 val df2 = df1 .groupBy("group_key") // 仅按group_key分组,这才是正确的分组逻辑 .agg(similarRowUdaf.alias("filtered_rows")) // 如果需要跨分区按某列取TopN,展开列表后再处理 import org.apache.spark.sql.functions.explode df2 .select($"group_key", explode($"filtered_rows").alias("row")) .select($"group_key", $"row.*") .orderBy($"group_key", $"length".desc) // 按length列降序排序 .groupBy("group_key") .agg(collect_list(struct(df1.columns.map(col): _*)).alias("top_rows")) // 收集TopN行 .show()
关键提示
- 替换筛选逻辑:把
mergeValue和mergeCombiners里的示例去重逻辑,换成你自己的相似行判断逻辑(比如比较特定列的相似度)。 - 控制分区数据量:通过
limit参数限制每个分区内每组的行数量,避免分区数据失衡,减轻跨分区合并的压力。 - 跨分区TopN:聚合后用
explode展开列表,按目标列排序后重新分组收集,就能实现跨分区的TopN需求。
内容的提问来源于stack exchange,提问作者Rubber Duck
相关产品推荐
相关产品推荐

