自定义函数

Spark 自定义函数:单行处理用 UDF,多行聚合用 UDAF

Spark SQL 内置了上百个函数——sum、avg、concat、when……但总有业务逻辑是内置函数覆盖不到的。

这时候就得上自定义函数:

类型 输入 → 输出 典型场景
UDF 一行 → 一行 字段转换、格式处理、数据清洗
UDAF 多行 → 一行 自定义聚合、加权平均、中位数

一句话:处理单条记录用 UDF,处理多行聚合用 UDAF

UDF:一行输入,一行输出

UDF 是”对每条记录单独处理”——输入一行的若干字段,输出一个值。

1.1 注册和使用 UDF

import org.apache.spark.sql.SparkSession

val spark = SparkSession.builder().appName("UDF").master("local[*]").getOrCreate()

// 1. 定义 UDF(把名字加上前缀)
val addPrefix = (name: String) => s"user_$name"

// 2. 注册 UDF
spark.udf.register("add_prefix", addPrefix)

// 3. 使用(SQL 或 DataFrame)
val df = Seq("Alice", "Bob").toDF("name")
df.createOrReplaceTempView("users")
spark.sql("SELECT name, add_prefix(name) AS username FROM users").show()
// +-----+----------+
// | name|  username|
// +-----+----------+
// |Alice|user_Alice|
// |  Bob|  user_Bob|
// +-----+----------+

1.2 UDF 支持多参数

// 两个参数:名字 + 前缀
val addPrefix2 = (name: String, prefix: String) => s"$prefix:$name"
spark.udf.register("add_prefix2", addPrefix2)

spark.sql("SELECT add_prefix2(name, 'emp') FROM users").show()
// +-------------------+
// |add_prefix2(name, emp)|
// +-------------------+
// |          emp:Alice|
// |            emp:Bob|
// +-------------------+

1.3 UDF 支持复杂类型

// 输入是数组,输出是数组长度
spark.udf.register("array_len", (arr: Array[Int]) => arr.length)
spark.sql("SELECT array_len(array(1,2,3,4))").show()   // 4

UDF 的限制: 只能一行一行地处理,不能跨行聚合。

UDAF:多行输入,一行输出

UDAF 是”多行数据聚合成一个值”——比如自定义平均值、加权平均、中位数。

Spark 2.x 之后推荐用 Aggregator 实现 UDAF。

2.1 实现一个自定义平均年龄

import org.apache.spark.sql.{Encoder, Encoders}
import org.apache.spark.sql.expressions.Aggregator

// 缓冲区:存"总年龄"和"总人数"
case class AgeBuffer(totalAge: Long, count: Long)

class AvgAge extends Aggregator[Long, AgeBuffer, Double] {
  // 1. 初始值
  override def zero: AgeBuffer = AgeBuffer(0L, 0L)

  // 2. 每条数据更新缓冲区
  override def reduce(buf: AgeBuffer, age: Long): AgeBuffer =
    AgeBuffer(buf.totalAge + age, buf.count + 1)

  // 3. 合并两个分区的缓冲区(分布式)
  override def merge(b1: AgeBuffer, b2: AgeBuffer): AgeBuffer =
    AgeBuffer(b1.totalAge + b2.totalAge, b1.count + b2.count)

  // 4. 计算最终结果
  override def finish(buf: AgeBuffer): Double =
    if (buf.count == 0) 0.0 else buf.totalAge.toDouble / buf.count

  // 5. 编码器
  override def bufferEncoder: Encoder[AgeBuffer] = Encoders.product
  override def outputEncoder: Encoder[Double] = Encoders.scalaDouble
}

2.2 注册和使用 UDAF

import org.apache.spark.sql.functions

// 注册 UDAF
spark.udf.register("avg_age", functions.udaf(new AvgAge()))

// 使用
val df = Seq(18, 20, 22).toDF("age")
df.createOrReplaceTempView("ages")
spark.sql("SELECT avg_age(age) FROM ages").show()   // 20.0

// DataFrame 方式也可以
import spark.implicits._
val avgCol = new AvgAge().toColumn.name("avg_age")
df.select(avgCol).show()

复杂 UDAF:加权平均值

case class Weighted(age: Long, weight: Double)
case class WeightedBuffer(sumAgeWeight: Double, sumWeight: Double)

class WeightedAvg extends Aggregator[Weighted, WeightedBuffer, Double] {
  override def zero: WeightedBuffer = WeightedBuffer(0.0, 0.0)

  override def reduce(buf: WeightedBuffer, input: Weighted): WeightedBuffer =
    WeightedBuffer(
      buf.sumAgeWeight + input.age * input.weight,
      buf.sumWeight + input.weight
    )

  override def merge(b1: WeightedBuffer, b2: WeightedBuffer): WeightedBuffer =
    WeightedBuffer(b1.sumAgeWeight + b2.sumAgeWeight, b1.sumWeight + b2.sumWeight)

  override def finish(buf: WeightedBuffer): Double =
    if (buf.sumWeight == 0) 0.0 else buf.sumAgeWeight / buf.sumWeight

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

// 使用
val data = Seq(Weighted(18, 0.5), Weighted(20, 0.5)).toDS()
data.select(new WeightedAvg().toColumn.name("weighted_avg")).show()  // 19.0

使用注意事项

1. UDF 性能比内置函数差

UDF 是黑盒,Catalyst 优化器没法优化它。能用内置函数就别用 UDF。

//  慢:自己写 UDF 转大写
spark.udf.register("my_upper", (s: String) => s.toUpperCase)
spark.sql("SELECT my_upper(name) FROM users")

//  快:用内置函数
spark.sql("SELECT upper(name) FROM users")   // 内置的,有优化

2. UDAF 的 merge 要正确实现

merge 是把两个分区的中间结果合并,要确保合并逻辑正确(通常是加法)。

3. 处理空值

输入可能为 null,UDF 里要处理:

val safeUpper = (s: String) => Option(s).map(_.toUpperCase).orNull
spark.udf.register("safe_upper", safeUpper)