SparkSQL自定义udf

自定义udf求平均值:

测试数据:

{"name":"Michael", "salary":3000}
{"name":"Andy", "salary":4500}
{"name":"Justin", "salary":3500}
{"name":"Berta", "salary":4000}

第一种方法:

def main(args: Array[String]): Unit = {
val spark = SparkSession.builder().appName("myUdf").master("local[*]").getOrCreate()
        spark.udf.register("myAverage", MyAverage)
        val df: DataFrame = spark.read.json("employees.json")
        df.createOrReplaceTempView("employees")
        val result = spark.sql("select myAverage(salary) as avg from employees")
        result.show()
}
object MyAverage extends UserDefinedAggregateFunction {
  override def inputSchema: StructType = StructType(StructField("salary", 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 = {
    if (!input.isNullAt(0)) {
      buffer(0) = buffer.getLong(0) + input.getLong(0)//将每一条salary进行累加
      buffer(1) = buffer.getLong(1) + 1//数量的累加
    }
  }
//合并
  override def merge(buffer1: MutableAggregationBuffer, buffer2: Row): Unit = {
    buffer1(0) = buffer1.getLong(0) + buffer2.getLong(0)//sum合并
    buffer1(1) = buffer1.getLong(1) + buffer2.getLong(1)//count合并
  }
//最后的处理
  override def evaluate(buffer: Row): Any = {
    buffer.getLong(0).toDouble / buffer.getLong(1)//求平均值
  }
}

第二种方法:

  def main(args: Array[String]): Unit = {
    val spark = SparkSession.builder().appName("myUdf").master("local[*]").getOrCreate()
        spark.udf.register("myAverage", MyAverage)
        val df: DataFrame = spark.read.json("employees.json")    
import spark.implicits._
        val ds = df.as[Employee]
        ds.show()
        val avg: TypedColumn[Employee, Double] = MyAverage1.toColumn.name("avg")
        ds.select(avg).show()
}
case class Employee(name: String, salary: Long)

case class Average(var sum: Long, var count: Long)

object MyAverage1 extends Aggregator[Employee, Average, Double] {
  override def zero: Average = Average(0L, 0L)

  override def reduce(b: Average, a: Employee): Average = {
    b.sum += a.salary
    b.count += 1
    b
  }

  override def merge(b1: Average, b2: Average): Average = {
    b1.sum += b2.sum
    b1.count += b2.count
    b1
  }

  override def finish(reduction: Average): Double = {
    reduction.sum.toDouble / reduction.count
  }

  override def bufferEncoder: Encoder[Average] = Encoders.product

  override def outputEncoder: Encoder[Double] = Encoders.scalaDouble
}

代码来自官网

评论 1
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

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

抵扣说明:

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

余额充值