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

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()

原代码的问题

  1. Aggregator实例化方式错了:Scala里Aggregator是抽象类,不能直接new并传参数,得继承它实现对应方法。
  2. createCombiner定义错误:它需要是把单个Row转成中间状态List[Row]的函数,不是直接写List[Row],应该写成(row: Row) => List(row)。
  3. mergeValue方法签名不对:这个方法是把单个输入行合并到缓冲区,参数应该是(buffer: List[Row], input: Row),不是两个List[Row]。
  4. 缺少Encoder定义:Aggregator必须为中间状态和输出类型提供Encoder,要么重写对应方法,要么用隐式上下文的Encoder。
  5. 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.26 06:15:39