在 Spark SQL 中,UDF(用户自定义函数)允许扩展 SQL 功能,实现自定义的业务逻辑。以下是完整的 UDF 使用指南。

1. 基础 UDF 定义和注册

标量 UDF(最常用)

import org.apache.spark.sql.functions.udf
import spark.implicits._

// 方法一:使用 udf() 函数包装 Scala 函数
val calculateBonus = (salary: Double, performance: String) => {
  performance match {
    case "excellent" => salary * 0.2
    case "good"      => salary * 0.1
    case _           => salary * 0.05
  }
}

// 注册 UDF
val bonusUDF = udf(calculateBonus)

// 方法二:直接定义并注册
spark.udf.register("calculate_bonus", calculateBonus)

在 DataFrame API 中使用 UDF

val employeesDF = Seq(
  ("Alice", 75000, "excellent"),
  ("Bob", 65000, "good"),
  ("Charlie", 80000, "average")
).toDF("name", "salary", "performance")

// 使用 UDF
val result = employeesDF
  .withColumn("bonus_amount", bonusUDF($"salary", $"performance"))
  .withColumn("total_compensation", $"salary" + $"bonus_amount")

result.show()

在 SQL 中使用 UDF

// 注册为 SQL 函数
spark.udf.register("sql_calculate_bonus", calculateBonus)

// 创建临时视图
employeesDF.createOrReplaceTempView("employees")

// 在 SQL 查询中使用
val sqlResult = spark.sql("""
  SELECT 
    name,
    salary,
    performance,
    sql_calculate_bonus(salary, performance) as bonus_amount,
    salary + sql_calculate_bonus(salary, performance) as total_compensation
  FROM employees
  WHERE performance IS NOT NULL
""")

sqlResult.show()

2. 高级 UDF 类型

处理复杂数据类型

// 处理数组类型的 UDF
val arrayStatsUDF = udf((arr: Seq[Int]) => {
  if (arr == null || arr.isEmpty) (0.0, 0.0, 0)
  else (arr.sum.toDouble, arr.sum.toDouble / arr.length, arr.length)
})

// 处理 Map 类型的 UDF  
val mapExtractUDF = udf((map: Map[String, String], key: String) => {
  map.getOrElse(key, "unknown")
})

// 使用示例
val complexDF = Seq(
  (Seq(1, 2, 3, 4, 5), Map("dept" -> "Engineering", "level" -> "Senior")),
  (Seq(10, 20), Map("dept" -> "Sales", "level" -> "Junior"))
).toDF("scores", "attributes")

val processed = complexDF
  .withColumn("array_stats", arrayStatsUDF($"scores"))
  .withColumn("department", mapExtractUDF($"attributes", lit("dept")))

processed.show(false)

返回结构体类型的 UDF

import org.apache.spark.sql.Row

// 返回多个值的 UDF
val analyzeTextUDF = udf((text: String) => {
  if (text == null) Row(0, 0, 0.0)
  else {
    val words = text.split("\\s+")
    val wordCount = words.length
    val charCount = text.length
    val avgWordLength = if (wordCount > 0) charCount.toDouble / wordCount else 0.0
    Row(wordCount, charCount, avgWordLength)
  }
})

// 定义返回类型
import org.apache.spark.sql.types._
val textStatsSchema = StructType(Array(
  StructField("word_count", IntegerType),
  StructField("char_count", IntegerType), 
  StructField("avg_word_length", DoubleType)
))

// 注册带 schema 的 UDF
spark.udf.register("analyze_text", analyzeTextUDF, textStatsSchema)

3. 性能优化技巧

避免 UDF 中的重复计算

// 不佳:每次调用都重新计算
val inefficientUDF = udf((x: Int) => {
  // 昂贵的计算
  Thread.sleep(100) // 模拟耗时操作
  x * x
})

// 优化:缓存计算结果或使用内置函数
val efficientUDF = udf((x: Int) => x * x) // 简单计算

// 或者尽可能使用 Spark 内置函数
// val squared = pow($"x", 2)  // 使用内置幂函数

使用 Pandas UDF(PySpark)

# PySpark 中的 Pandas UDF(性能更好)
from pyspark.sql.functions import pandas_udf
from pyspark.sql.types import DoubleType

@pandas_udf(DoubleType())
def pandas_calculate_bonus(salary_series, performance_series):
    # 向量化操作,性能远超逐行处理的 UDF
    bonus_rates = performance_series.map({
        'excellent': 0.2,
        'good': 0.1, 
        'average': 0.05
    }).fillna(0.05)
    
    return salary_series * bonus_rates

4. 错误处理和空值安全

安全的 UDF 设计

// 不安全版本
val unsafeUDF = udf((str: String) => str.length)

// 安全版本:处理 null 和异常
val safeUDF = udf((str: String) => {
  try {
    if (str == null) 0 else str.length
  } catch {
    case _: Exception => 0  // 处理任何异常
  }
})

// 更精细的错误处理
val robustUDF = udf((str: String) => {
  Option(str) match {
    case Some(s) if s.nonEmpty => s.length
    case Some(s) if s.isEmpty => 0
    case None => 0
    case _ => -1  // 未知情况
  }
})

5. 实际业务场景示例

场景1:数据清洗和标准化

// 电话号码格式化 UDF
val formatPhoneUDF = udf((phone: String) => {
  if (phone == null) null
  else {
    val digits = phone.replaceAll("\\D", "")  // 移除非数字字符
    if (digits.length == 10) s"(${digits.substring(0,3)}) ${digits.substring(3,6)}-${digits.substring(6)}"
    else if (digits.length == 11 && digits.startsWith("1")) s"1 (${digits.substring(1,4)}) ${digits.substring(4,7)}-${digits.substring(7)}"
    else phone  // 无法格式化,返回原值
  }
})

// 邮箱验证 UDF
val isValidEmailUDF = udf((email: String) => {
  if (email == null) false
  else {
    val emailRegex = """^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$""".r
    emailRegex.findFirstIn(email).isDefined
  }
})

场景2:业务逻辑计算

// 客户等级计算 UDF
val calculateCustomerTierUDF = udf((totalSpent: Double, orderCount: Int, membershipYears: Int) => {
  val score = totalSpent * 0.5 + orderCount * 10 + membershipYears * 50
  
  if (score >= 10000) "Platinum"
  else if (score >= 5000) "Gold" 
  else if (score >= 1000) "Silver"
  else "Bronze"
})

// 折扣计算 UDF
val calculateDiscountUDF = udf((originalPrice: Double, customerTier: String, couponCode: String) => {
  val tierDiscount = customerTier match {
    case "Platinum" => 0.15
    case "Gold"     => 0.10
    case "Silver"   => 0.05
    case _          => 0.0
  }
  
  val couponDiscount = couponCode match {
    case "SAVE20" => 0.20
    case "SAVE10" => 0.10
    case "WELCOME" => 0.05
    case _ => 0.0
  }
  
  val totalDiscount = Math.min(tierDiscount + couponDiscount, 0.30)  // 最大折扣30%
  originalPrice * totalDiscount
})

场景3:地理空间计算

// 距离计算 UDF(简化版哈弗辛公式)
val calculateDistanceUDF = udf((lat1: Double, lon1: Double, lat2: Double, lon2: Double) => {
  val earthRadius = 6371.0  // 地球半径,公里
  
  val dLat = Math.toRadians(lat2 - lat1)
  val dLon = Math.toRadians(lon2 - lon1)
  
  val a = Math.sin(dLat/2) * Math.sin(dLat/2) +
          Math.cos(Math.toRadians(lat1)) * Math.cos(Math.toRadians(lat2)) *
          Math.sin(dLon/2) * Math.sin(dLon/2)
  
  val c = 2 * Math.atan2(Math.sqrt(a), Math.sqrt(1-a))
  earthRadius * c
})

6. UDF 管理和最佳实践

UDF 注册管理器

class UDFManager(spark: SparkSession) {
  private val registeredUDFs = scala.collection.mutable.Set[String]()
  
  def registerScalarUDF(name: String, function: AnyRef, dataType: DataType): Unit = {
    spark.udf.register(name, function)
    registeredUDFs.add(name)
    println(s"已注册标量 UDF: $name")
  }
  
  def listRegisteredUDFs(): Unit = {
    println("已注册的 UDF:")
    registeredUDFs.foreach(println)
  }
  
  def unregisterUDF(name: String): Unit = {
    // Spark 不支持直接注销 UDF,但可以从内部集合移除
    registeredUDFs.remove(name)
    println(s"已从管理器移除 UDF: $name")
  }
}

// 使用示例
val udfManager = new UDFManager(spark)
udfManager.registerScalarUDF("format_phone", formatPhoneUDF, StringType)

性能测试工具

def benchmarkUDF(udfName: String, testDF: DataFrame, iterations: Int = 10): Unit = {
  val startTime = System.currentTimeMillis()
  
  for (i <- 1 to iterations) {
    testDF.selectExpr(s"$udfName(column)").collect()
  }
  
  val endTime = System.currentTimeMillis()
  val avgTime = (endTime - startTime).toDouble / iterations
  
  println(s"UDF '$udfName' 平均执行时间: ${avgTime}ms (${iterations}次迭代)")
}

最佳实践总结

  1. 优先使用内置函数:Spark 内置函数通常比 UDF 性能更好
  2. 处理空值和异常:确保 UDF 能够优雅处理边界情况
  3. 保持 UDF 简单:复杂的业务逻辑尽量在数据管道的前端处理
  4. 测试性能影响:在生产环境使用前进行性能测试
  5. 文档化 UDF:为每个 UDF 提供清晰的文档说明用途和参数

通过合理使用 UDF,可以极大地扩展 Spark SQL 的功能,实现复杂的业务逻辑,同时保持代码的可维护性和性能。

Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐