在 Spark SQL 中,UDF(用户自定义函数)是扩展 SQL 功能的重要方式。以下是完整的 UDF 使用指南。

1. UDF 基础概念

1.1 UDF 类型

  • 普通 UDF:一对一转换,输入一行输出一行
  • UDAF:用户自定义聚合函数,多行输入一行输出
  • UDTF:用户自定义表生成函数,一行输入多行输出

2. 基本 UDF 实现

2.1 字符串处理 UDF

from pyspark.sql import SparkSession
from pyspark.sql.functions import udf
from pyspark.sql.types import StringType, IntegerType
import re

# 创建 SparkSession
spark = SparkSession.builder \
    .appName("UDF Demo") \
    .config("spark.sql.adaptive.enabled", "true") \
    .getOrCreate()

# 示例数据
data = [("Alice", 25, "alice@email.com"), 
        ("Bob", 30, "bob.company@test.com"),
        ("Charlie", 35, "charlie123@gmail.com")]

df = spark.createDataFrame(data, ["name", "age", "email"])
df.show()

2.2 注册和使用 UDF

# 方法1:使用装饰器注册 UDF
@udf(returnType=StringType())
def extract_domain(email):
    """提取邮箱域名"""
    if email and '@' in email:
        return email.split('@')[1]
    return None

# 方法2:直接注册 UDF
def calculate_age_group(age):
    """根据年龄分组"""
    if age < 18:
        return "未成年"
    elif age < 30:
        return "青年"
    elif age < 50:
        return "中年"
    else:
        return "老年"

age_group_udf = udf(calculate_age_group, StringType())

# 在 DataFrame 中使用 UDF
result_df = df.select(
    "name",
    "age", 
    "email",
    extract_domain("email").alias("domain"),
    age_group_udf("age").alias("age_group")
)

result_df.show()

3. 复杂数据类型 UDF

3.1 处理数组和映射类型

from pyspark.sql.types import ArrayType, MapType, StructType, StructField
import json

# 示例数据:包含数组和JSON字符串
complex_data = [
    ("Alice", "[1, 2, 3, 4, 5]", '{"city": "Beijing", "score": 85}'),
    ("Bob", "[10, 20, 30]", '{"city": "Shanghai", "score": 92}'),
    ("Charlie", "[7, 8, 9, 10]", '{"city": "Guangzhou", "score": 78}')
]

complex_df = spark.createDataFrame(complex_data, ["name", "numbers", "profile"])

# 数组处理的 UDF
@udf(returnType=IntegerType())
def calculate_average(numbers_str):
    """计算数组平均值"""
    try:
        numbers = json.loads(numbers_str)
        return sum(numbers) // len(numbers) if numbers else 0
    except:
        return 0

# JSON 处理的 UDF  
@udf(returnType=StringType())
def extract_city(profile_str):
    """从JSON中提取城市"""
    try:
        profile = json.loads(profile_str)
        return profile.get('city', '未知')
    except:
        return '解析错误'

# 使用 UDF
complex_result = complex_df.select(
    "name",
    "numbers",
    "profile",
    calculate_average("numbers").alias("avg_number"),
    extract_city("profile").alias("city")
)

complex_result.show()

4. 多参数 UDF

4.1 接受多个输入的 UDF

from pyspark.sql.types import DoubleType

# 业务数据
business_data = [
    ("Product_A", 100, 10.5),  # 产品名,销量,单价
    ("Product_B", 250, 8.2),
    ("Product_C", 80, 15.0),
    ("Product_D", 300, 5.5)
]

business_df = spark.createDataFrame(business_data, ["product", "quantity", "price"])

# 多参数 UDF - 计算销售额和折扣
@udf(returnType=DoubleType())
def calculate_revenue(quantity, price, discount_rate=0.1):
    """计算销售额(含折扣)"""
    revenue = quantity * price
    discounted_revenue = revenue * (1 - discount_rate)
    return round(discounted_revenue, 2)

@udf(returnType=StringType())  
def get_sales_category(quantity, revenue):
    """根据销量和销售额分类"""
    if quantity > 200 and revenue > 1500:
        return "热销产品"
    elif quantity > 100 and revenue > 800:
        return "畅销产品" 
    elif quantity > 50:
        return "普通产品"
    else:
        return "滞销产品"

# 使用多参数 UDF
business_result = business_df.select(
    "product",
    "quantity",
    "price",
    calculate_revenue("quantity", "price").alias("revenue"),
    calculate_revenue("quantity", "price", udf(lambda: 0.2, DoubleType())()).alias("revenue_20p_discount")
)

# 添加分类(需要先计算revenue)
business_result = business_result.withColumn(
    "category", 
    get_sales_category("quantity", "revenue")
)

business_result.show()

5. 在 Spark SQL 中使用 UDF

5.1 注册 UDF 并在 SQL 中使用

# 注册 UDF 到 Spark SQL
spark.udf.register("sql_extract_domain", extract_domain)
spark.udf.register("sql_age_group", calculate_age_group)
spark.udf.register("sql_calculate_revenue", calculate_revenue)

# 创建临时视图
df.createOrReplaceTempView("people")
business_df.createOrReplaceTempView("products")

# 在 SQL 中使用 UDF
sql_result = spark.sql("""
    SELECT 
        name,
        age,
        email,
        sql_extract_domain(email) as domain,
        sql_age_group(age) as age_group
    FROM people
    WHERE sql_age_group(age) IN ('青年', '中年')
""")

sql_result.show()

# 复杂的 SQL 查询 with UDF
business_sql = spark.sql("""
    SELECT 
        product,
        quantity,
        price,
        sql_calculate_revenue(quantity, price) as revenue,
        CASE 
            WHEN quantity > 200 AND sql_calculate_revenue(quantity, price) > 1500 THEN '热销'
            WHEN quantity > 100 AND sql_calculate_revenue(quantity, price) > 800 THEN '畅销'
            ELSE '普通'
        END as sales_category
    FROM products
    ORDER BY revenue DESC
""")

business_sql.show()

6. 高级 UDF 技巧

6.1 使用外部库的 UDF

import phonenumbers
from datetime import datetime

# 电话号码验证 UDF
@udf(returnType=StringType())
def validate_phone_number(phone, country_code="CN"):
    """验证电话号码格式"""
    try:
        parsed_number = phonenumbers.parse(phone, country_code)
        if phonenumbers.is_valid_number(parsed_number):
            return phonenumbers.format_number(parsed_number, 
                                            phonenumbers.PhoneNumberFormat.INTERNATIONAL)
        else:
            return "无效号码"
    except:
        return "格式错误"

# 日期处理 UDF
@udf(returnType=StringType())
def format_date(date_str, input_format="%Y-%m-%d", output_format="%Y年%m月%d日"):
    """日期格式转换"""
    try:
        date_obj = datetime.strptime(date_str, input_format)
        return date_obj.strftime(output_format)
    except:
        return date_str  # 解析失败返回原值

# 测试数据
contact_data = [
    ("Alice", "13800138000", "2023-01-15"),
    ("Bob", "+86-13900139000", "2023-02-20"), 
    ("Charlie", "12345", "2023-03-25")  # 无效号码
]

contact_df = spark.createDataFrame(contact_data, ["name", "phone", "date"])

contact_result = contact_df.select(
    "name",
    "phone",
    "date",
    validate_phone_number("phone").alias("formatted_phone"),
    format_date("date").alias("chinese_date")
)

contact_result.show()

6.2 性能优化的 UDF

from pyspark.sql.types import BooleanType
import pandas as pd

# 避免在 UDF 中创建大量小对象
# 不好的写法(每次调用都创建新对象)
@udf(returnType=BooleanType())
def is_valid_email_slow(email):
    """验证邮箱格式(性能较差)"""
    import re
    pattern = re.compile(r'^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$')
    return bool(pattern.match(email)) if email else False

# 好的写法(预编译正则表达式)
_email_pattern = None

def get_email_pattern():
    """延迟加载并缓存正则表达式"""
    global _email_pattern
    if _email_pattern is None:
        import re
        _email_pattern = re.compile(r'^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$')
    return _email_pattern

@udf(returnType=BooleanType())
def is_valid_email_fast(email):
    """验证邮箱格式(性能优化)"""
    pattern = get_email_pattern()
    return bool(pattern.match(email)) if email else False

# 测试性能
test_data = [("test@example.com",)] * 10000
test_df = spark.createDataFrame(test_data, ["email"])

# 使用优化后的 UDF
result_fast = test_df.filter(is_valid_email_fast("email"))
print(f"优化版过滤后行数: {result_fast.count()}")

7. Pandas UDF(向量化 UDF)

7.1 使用 Pandas UDF 提升性能

from pyspark.sql.functions import pandas_udf
import pandas as pd

# 示例数据:大规模数值计算
large_data = [(i, i * 10, i * 0.5) for i in range(1, 10001)]
large_df = spark.createDataFrame(large_data, ["id", "value1", "value2"])

# 普通的 Python UDF(逐行处理)
@udf(returnType=DoubleType())
def calculate_complex_value_regular(value1, value2):
    """复杂计算 - 普通UDF"""
    import math
    return math.log(value1 + 1) * math.sqrt(value2 + 1)

# Pandas UDF(向量化处理)
@pandas_udf(returnType=DoubleType())
def calculate_complex_value_pandas(value1: pd.Series, value2: pd.Series) -> pd.Series:
    """复杂计算 - Pandas UDF(向量化)"""
    import numpy as np
    return np.log(value1 + 1) * np.sqrt(value2 + 1)

# 性能对比
import time

# 普通 UDF
start_time = time.time()
regular_result = large_df.withColumn("result_regular", 
                                   calculate_complex_value_regular("value1", "value2"))
regular_result.count()  # 触发计算
regular_time = time.time() - start_time

# Pandas UDF  
start_time = time.time()
pandas_result = large_df.withColumn("result_pandas", 
                                  calculate_complex_value_pandas("value1", "value2"))
pandas_result.count()  # 触发计算
pandas_time = time.time() - start_time

print(f"普通UDF执行时间: {regular_time:.2f}秒")
print(f"Pandas UDF执行时间: {pandas_time:.2f}秒")
print(f"性能提升: {regular_time/pandas_time:.1f}倍")

# 显示部分结果对比
comparison = pandas_result.select("id", "value1", "value2", 
                                "result_regular", "result_pandas").limit(10)
comparison.show()

8. 错误处理和调试

8.1 UDF 错误处理最佳实践

from pyspark.sql.types import StringType

# 安全的 UDF 实现
@udf(returnType=StringType())
def safe_string_operation(text, operation_type="upper"):
    """安全的字符串操作 UDF"""
    try:
        if not text:
            return ""
            
        if operation_type == "upper":
            return text.upper()
        elif operation_type == "lower":
            return text.lower()
        elif operation_type == "length":
            return str(len(text))
        else:
            return text
            
    except Exception as e:
        # 记录错误但不中断处理
        print(f"UDF执行错误: {e}")
        return f"错误: {str(e)}"

# 测试错误处理
error_test_data = [
    ("hello", "upper"),
    (None, "upper"),  # None值
    (123, "upper"),   # 错误类型
    ("world", "invalid_operation")  # 无效操作
]

error_df = spark.createDataFrame(error_test_data, ["text", "operation"])

error_result = error_df.select(
    "text",
    "operation", 
    safe_string_operation("text", "operation").alias("result")
)

error_result.show()

9. UDF 管理和最佳实践

9.1 UDF 注册和管理工具类

class UDFManager:
    """UDF 管理类"""
    
    def __init__(self, spark_session):
        self.spark = spark_session
        self.registered_udfs = {}
    
    def register_udf(self, name, func, return_type):
        """注册 UDF 并记录元数据"""
        udf_function = udf(func, return_type)
        self.spark.udf.register(name, udf_function)
        self.registered_udfs[name] = {
            'function': func,
            'return_type': return_type,
            'description': func.__doc__ or 'No description'
        }
        print(f"UDF '{name}' 注册成功")
    
    def list_udfs(self):
        """列出所有已注册的 UDF"""
        print("已注册的 UDF:")
        for name, info in self.registered_udfs.items():
            print(f"- {name}: {info['description']}")
    
    def get_udf_info(self, name):
        """获取 UDF 详细信息"""
        if name in self.registered_udfs:
            info = self.registered_udfs[name]
            print(f"UDF: {name}")
            print(f"描述: {info['description']}")
            print(f"返回类型: {info['return_type']}")
        else:
            print(f"UDF '{name}' 未找到")

# 使用 UDF 管理器
udf_manager = UDFManager(spark)

# 注册多个 UDF
udf_manager.register_udf("extract_domain", extract_domain, StringType())
udf_manager.register_udf("age_group", calculate_age_group, StringType())
udf_manager.register_udf("calculate_revenue", calculate_revenue, DoubleType())

# 查看注册的 UDF
udf_manager.list_udfs()
udf_manager.get_udf_info("extract_domain")

10. 实际业务场景示例

10.1 电商数据分析 UDF

# 电商业务相关的 UDF
from pyspark.sql.types import TimestampType
from datetime import datetime, timedelta

@udf(returnType=StringType())
def categorize_price(price):
    """价格区间分类"""
    if price < 50:
        return "低价"
    elif price < 200:
        return "中价" 
    elif price < 500:
        return "高价"
    else:
        return "奢侈价"

@udf(returnType=StringType())
def calculate_delivery_time(order_time, delivery_time):
    """计算配送时效"""
    try:
        if order_time and delivery_time:
            order_dt = datetime.fromtimestamp(order_time / 1000)  # 假设是毫秒时间戳
            delivery_dt = datetime.fromtimestamp(delivery_time / 1000)
            duration = delivery_dt - order_dt
            hours = duration.total_seconds() / 3600
            
            if hours < 24:
                return "当日达"
            elif hours < 48:
                return "次日达"
            else:
                return f"{int(hours//24)}天达"
    except:
        pass
    return "未知"

@udf(returnType=StringType())
def generate_recommendation(category, price, sales_count):
    """生成商品推荐标签"""
    tags = []
    
    if sales_count > 1000:
        tags.append("爆款")
    elif sales_count > 100:
        tags.append("热销")
        
    if price < 100:
        tags.append("性价比高")
        
    if category in ["电子", "数码"]:
        tags.append("科技新品")
    elif category in ["服装", "时尚"]:
        tags.append("潮流单品")
        
    return " | ".join(tags) if tags else "普通商品"

# 模拟电商数据
ecommerce_data = [
    ("手机", 2999, 1500, 1672531200000, 1672560000000),  # 商品,价格,销量,下单时间,送达时间
    ("书籍", 45, 80, 1672531200000, 1672617600000),
    ("服装", 199, 300, 1672531200000, 1672574400000),
    ("家电", 1299, 50, 1672531200000, 1672704000000)
]

ecommerce_df = spark.createDataFrame(ecommerce_data, 
                                   ["category", "price", "sales_count", 
                                    "order_time", "delivery_time"])

# 应用业务 UDF
ecommerce_analysis = ecommerce_df.select(
    "category",
    "price", 
    "sales_count",
    categorize_price("price").alias("price_category"),
    calculate_delivery_time("order_time", "delivery_time").alias("delivery_speed"),
    generate_recommendation("category", "price", "sales_count").alias("recommendation_tags")
)

ecommerce_analysis.show(truncate=False)

11. 总结

Spark SQL UDF 核心要点:

  1. 注册方式udf() 函数或 spark.udf.register()
  2. 类型安全:必须指定返回类型
  3. 性能优化:使用 Pandas UDF 进行向量化计算
  4. 错误处理:在 UDF 内部处理异常
  5. SQL 集成:注册后可在 Spark SQL 中直接使用

最佳实践:

  • 为复杂的业务逻辑创建专用的 UDF
  • 使用 Pandas UDF 处理大规模数据
  • 在 UDF 中包含适当的错误处理
  • 为 UDF 添加文档字符串说明用途
  • 考虑性能影响,避免在 UDF 中执行重操作

通过 UDF,你可以极大地扩展 Spark SQL 的功能,使其能够处理各种复杂的业务逻辑和数据转换需求。

Logo

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

更多推荐