1.基于UserDefinedAggregateFunction实现平均数的计算
package com.bigdata.wb.spark
import org.apache.spark.sql.Row
import org.apache.spark.sql.expressions.{MutableAggregationBuffer, UserDefinedAggregateFunction}
import org.apache.spark.sql.types._
object SparkUDAF extends UserDefinedAggregateFunction{
//输入数据类型
override def inputSchema: StructType = StructType(StructField("input", LongType)::Nil)
//缓冲区数据类型
override def bufferSchema: StructType = StructType(StructField("sum", LongType)::StructField("count", LongType)::Nil)
//聚合之后输出数据类型
override def dataType: DataType = DoubleType
//相同输入是否总能得到相同输出
override def deterministic: Boolean = true
//初始化缓冲区
override def initialize(buffer: MutableAggregationBuffer): Unit = {
buffer(0) = 0L
buffer(1) = 0L
}
//给聚合函数传入一条数据进行处理
override def update(buffer: MutableAggregationBuffer, input: Row): Unit = {
buffer(0) = buffer.getLong(0) + input.getLong(0)
buffer(1) = buffer.getLong(1) + 1
buffer.update(0, buffer(0))
buffer.update(1, buffer(1))
}
//合并聚合函数缓冲区
override def merge(buffer1: MutableAggregationBuffer, buffer2: Row): Unit = {
buffer1(0) = buffer1.getLong(0) + buffer2.getLong(0)
buffer1(1) = buffer1.getLong(1) + buffer2.getLong(1)
}
//计算最终结果
override def evaluate(buffer: Row): Any = buffer.getLong(0).toDouble / buffer.getLong(1)
}
2.基于Aggregator实现平均数的计算
package com.bigdata.wb.spark
import org.apache.spark.SparkConf
import org.apache.spark.sql.{Encoders, SparkSession, TypedColumn}
import org.apache.spark.sql.expressions.{Aggregator, UserDefinedAggregateFunction}
/**
* @ author spencer
* @ date 2020/7/14 13:46
*
* spark中UDAF中统计平均数
*/
case class Employee(name: String, salary: Long)
case class Average(var sum: Long, var count: Long)
object SparkUDAFDemo02{
def main(args: Array[String]): Unit = {
val conf = new SparkConf().setAppName("SparkSQLDemo").setMaster("local[*]")
val spark = SparkSession.builder()
.config(conf)
.enableHiveSupport()
.getOrCreate()
val employeeRDD = spark.read.json("file:///D:\\spark-2.3.0-bin-hadoop2.7\\examples\\src\\main\\resources\\employees.json")
import spark.implicits._
val employeeDF = employeeRDD.as[Employee]
//第一种方式调用udaf,spark3.0.0官网使用的这种方式
val averageSalary = MyUDAF.toColumn.name("average_salary")
employeeDF.select(averageSalary).show()
//第二种方式调用udaf
employeeDF.createOrReplaceTempView("employee")
spark.udf.register("averageSalary", SparkUDAF)
spark.sql("select averageSalary(salary) avg_salary from employee").show()
spark.stop()
}
}
/**
* spark3.0.0中自定义UDAF继承Aggregator实现
*/
object MyUDAF extends Aggregator[Employee, Average, Double]{
override def zero = Average(0L, 0L)
override def reduce(buffer: Average, employee: Employee) = {
buffer.sum += employee.salary
buffer.count += 1
buffer
}
override def merge(buffer1: Average, buffer2: Average) = {
buffer1.sum += buffer2.sum
buffer1.count += buffer2.count
buffer1
}
override def finish(reduction: Average) = reduction.sum / reduction.count
override def bufferEncoder = Encoders.product
override def outputEncoder = Encoders.scalaDouble
}
结果如下:
+--------------+
|average_salary|
+--------------+
| 3750.0|
+--------------+
+----------+
|avg_salary|
+----------+
| 3750.0|
+----------+