基于Scala实现返回Array类型的Spark UDAF:按时间排序渠道
实现按Time升序排序Channel的Spark UDAF(Scala)
没问题,我帮你搞定这个需求——我们需要写一个自定义聚合函数(UDAF),按id分组后,把channel字段按照time的升序排列,最终返回Array[String]类型,而且支持Spark SQL直接调用。
步骤1:编写UDAF类
我们继承UserDefinedAggregateFunction(这个类在SQL场景下适配性更好),核心逻辑是在聚合过程中收集每个分组的channel和time对,最后排序并提取channel组成数组:
import org.apache.spark.sql.expressions.UserDefinedAggregateFunction import org.apache.spark.sql.types._ import org.apache.spark.sql.{Row, SparkSession} class SortChannelsByTimeUDAF extends UserDefinedAggregateFunction { // 输入数据的结构:(channel字符串, time整数) override def inputSchema: StructType = StructType( StructField("channel", StringType) :: StructField("time", IntegerType) :: Nil ) // 缓冲区的结构:存储收集到的(channel, time)元组数组 override def bufferSchema: StructType = StructType( StructField("pairs", ArrayType(StructType( StructField("channel", StringType) :: StructField("time", IntegerType) :: Nil ))) :: Nil ) // 返回值类型:字符串数组 override def dataType: DataType = ArrayType(StringType) // 确定性标记:相同输入必然返回相同输出 override def deterministic: Boolean = true // 初始化缓冲区:空数组 override def initialize(buffer: Row): Unit = { buffer.update(0, Array.empty[(String, Int)]) } // 更新缓冲区:把当前输入的(channel, time)添加到缓冲区数组 override def update(buffer: Row, input: Row): Unit = { val currentPairs = buffer.getAs[Array[(String, Int)]](0) val newPair = (input.getAs[String](0), input.getAs[Int](1)) buffer.update(0, currentPairs :+ newPair) } // 合并缓冲区:把两个分区的缓冲区数组合并 override def merge(buffer1: Row, buffer2: Row): Unit = { val pairs1 = buffer1.getAs[Array[(String, Int)]](0) val pairs2 = buffer2.getAs[Array[(String, Int)]](0) buffer1.update(0, pairs1 ++ pairs2) } // 计算最终结果:按time升序排序,提取channel组成数组 override def evaluate(buffer: Row): Any = { val pairs = buffer.getAs[Array[(String, Int)]](0) pairs.sortBy(_._2).map(_._1) } }
步骤2:注册UDAF并执行SQL查询
接下来把UDAF注册到SparkSession,然后用SQL完成分组排序:
// 初始化SparkSession(本地测试用,生产环境去掉master配置) val spark = SparkSession.builder() .appName("SortChannelsUDAFTest") .master("local[*]") .getOrCreate() // 创建你的测试DataFrame val myDF = Seq( (1,"A",100), (1,"E",300), (1,"B",200), (2,"A",200), (2,"C",300), (2,"D",100) ).toDF("id","channel","time") // 注册自定义聚合函数 spark.udf.register("sort_channels_by_time", new SortChannelsByTimeUDAF()) // 执行Spark SQL查询 spark.sql(""" SELECT id, sort_channels_by_time(channel, time) AS sorted_channels FROM myDF GROUP BY id """).show(false)
预期输出
运行后会得到如下结果:
+---+----------------+ |id |sorted_channels | +---+----------------+ |1 |[A, B, E] | |2 |[D, A, C] | +---+----------------+
补充说明
- 如果你的Spark版本是3.0+,也可以用
Aggregator实现类型更安全的UDAF,但上面的UserDefinedAggregateFunction在SQL中调用更直接。 - 缓冲区用
Array而非ListBuffer是因为Spark的Row只能存储序列化后的集合类型,Array更适配Spark内部的序列化机制。
内容的提问来源于stack exchange,提问作者Regressor
相关产品推荐
相关产品推荐

