自定义函数
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)