在 Spark SQL 中,如何通过 UDF(用户自定义函数)扩展 SQL 功能?
·
在 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 核心要点:
- 注册方式:
udf()函数或spark.udf.register() - 类型安全:必须指定返回类型
- 性能优化:使用 Pandas UDF 进行向量化计算
- 错误处理:在 UDF 内部处理异常
- SQL 集成:注册后可在 Spark SQL 中直接使用
最佳实践:
- 为复杂的业务逻辑创建专用的 UDF
- 使用 Pandas UDF 处理大规模数据
- 在 UDF 中包含适当的错误处理
- 为 UDF 添加文档字符串说明用途
- 考虑性能影响,避免在 UDF 中执行重操作
通过 UDF,你可以极大地扩展 Spark SQL 的功能,使其能够处理各种复杂的业务逻辑和数据转换需求。
更多推荐


所有评论(0)