如何在Spark SQL中使用以Row为输入的Aggregator?
问题描述
我定义了一个以Row为输入的Spark Aggregator,用于统计行数:
val testAgg = new Aggregator[Row, Int, Int] { def zero = 0 def reduce(buffer: Int, row: Row):Int = buffer + 1 def merge(b1: Int, b2: Int):Int = b1 + b2 def finish(b: Int): Int = b def bufferEncoder: Encoder[Int] = Encoders.scalaInt def outputEncoder: Encoder[Int] = Encoders.scalaInt }
我知道如何在DataFrame API中使用这个聚合器,但不清楚如何在Spark SQL中调用它。我尝试了以下代码:
spark.udf.register("testAgg", functions.udaf(testAgg)) spark.sql("SELECT testAgg(*) from mytable2").show() // 无法运行 spark.sql("SELECT testAgg(struct(*)) from mytable2").show() // 同样无法运行
但报错:
Error: No applicable constructor/method found for zero actual parameters; candidates are: "public org.apache.spark.sql.Row org.apache.spark.sql.Row$.apply(scala.collection.Seq)"
请问是否可以在Spark SQL中使用以Row为输入的Aggregator?
解决方案
可以在Spark SQL中使用以Row为输入的Aggregator,问题出在两个核心点:注册UDAF时缺少隐式的Encoder[Row],以及SQL调用时的参数传递方式错误。
1. 补充隐式Encoder[Row]
注册UDAF时,Spark需要隐式的Encoder[Row]来处理输入类型,你需要在代码中添加对应的隐式编码器:
import org.apache.spark.sql.Encoders // 添加隐式Row编码器 implicit val rowEncoder: Encoder[Row] = Encoders.row
2. 正确的SQL调用方式
不能直接使用testAgg(*)——这会把每一列作为单独参数传递,而你的Aggregator期望单个Row参数,必须用struct(*)把所有列打包成一个结构体后再传递:
spark.sql("SELECT testAgg(struct(*)) from mytable2").show()
完整可运行示例
import org.apache.spark.sql.{Encoders, Row, SparkSession} import org.apache.spark.sql.expressions.Aggregator import org.apache.spark.sql.functions object RowAggregatorSQLDemo { def main(args: Array[String]): Unit = { val spark = SparkSession.builder() .appName("RowAggregatorSQL") .master("local[*]") .getOrCreate() import spark.implicits._ // 定义以Row为输入的Aggregator val testAgg = new Aggregator[Row, Int, Int] { def zero = 0 def reduce(buffer: Int, row: Row): Int = buffer + 1 def merge(b1: Int, b2: Int): Int = b1 + b2 def finish(b: Int): Int = b def bufferEncoder: Encoder[Int] = Encoders.scalaInt def outputEncoder: Encoder[Int] = Encoders.scalaInt } // 添加隐式Row编码器 implicit val rowEncoder: Encoder[Row] = Encoders.row // 注册UDAF spark.udf.register("testAgg", functions.udaf(testAgg)) // 创建测试表 val df = Seq((1, "a"), (2, "b"), (3, "c")).toDF("id", "name") df.createOrReplaceTempView("mytable2") // SQL中调用UDAF统计行数 spark.sql("SELECT testAgg(struct(*)) as total_rows from mytable2").show() } }
运行后会输出表的总行数,与count(*)结果一致。
内容的提问来源于stack exchange,提问作者yalei du
相关产品推荐
相关产品推荐

