spark中自定义UDAF函数实现的两种方式---UserDefinedAggregateFunction和Aggregator

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|
+----------+
  • 1
    点赞
  • 3
    收藏
    觉得还不错? 一键收藏
  • 0
    评论
评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值