Spark教學(四) - 萬用瑞士刀之UDF/UDTF

目錄

前言

資訊
本篇文章接續前篇 Spark教學(三) - Spark DataFrame 轉換 Pandas DataFrame(連結請依實際路徑調整),繼續往下編寫處理邏輯。如果手上還沒有可用的環境,建議先依照前面幾篇的步驟完成安裝,跟著本篇操作會更順暢。

UDF/UDTF

在探索了一段時間之後,相信大家會發現不少內建工具已經能應付日常的資料轉換需求,不論是 Spark DataFrame 還是 pandas API on Spark 的 DataFrame,都能解決大多數場景。但如果我們想要像寫一般程式那樣,自己寫一個 function、丟入參數、跑完之後拿到想要的結果,這件事在 Spark 裡可行嗎?

答案是可行的,而且 Spark 提供了對應的模組來實現,它叫做 UDF / UDTF。它們做的事情,簡單說就是把資料一列一列丟進 function 裡做轉換,最後輸出成一個 DataFrame。雖然效能上會比內建的原生函數慢一些,但換來的是完全依照使用者需求客製化的彈性,這一點相當實用。

那麼,UDF 與 UDTF 之間有什麼差別呢?我們先來看看下面的比較表。

Spark UDF 與 UDTF 差別解析

UDF (User-Defined Function)UDTF (User-Defined Table Function) 最核心的差異,在於「輸入對輸出的映射關係」與「回傳的資料結構」:

  • UDF1 對 1:傳入一列資料,吐出一個欄位值(純量 / Scalar Value)。
  • UDTF1 對多:傳入一列資料,透過 yield 展開成零到多列、多欄位的完整表格(Table)。
注意
值得一提的是,Python UDTF 是 Spark 3.5 才正式加入的新功能,跟很早就存在的 UDF 並不是同期的產物。如果你的叢集版本較舊,可能得先確認一下是否支援。

📊 UDF vs UDTF 詳細對比表

比較項目 UDF (User-Defined Function) UDTF (User-Defined Table Function)
映射關係 1 對 1(輸入 1 行 → 輸出 1 個值) 1 對多(輸入 1 行 → 輸出多列多欄表格)
回傳格式 (Return Type) 單一型別(例如 StringType()IntegerType() 表格結構(需定義多個欄位名稱與型別,如 "key: string, value: string"
PySpark 結構 @udf 裝飾器搭配一般 Python 函數 (Function) @udtf 裝飾器搭配包含 eval() 的 Python 類別 (Class)
回傳機制 使用 return 回傳結果 使用 yield 動態產出多筆資料
SQL 呼叫位置 放在 SELECT 欄位清單或 WHERE 條件中 放在 FROM 子句或搭配 LATERAL 查詢
典型應用場景 • 字串轉換與清理(如:去空白、轉大寫)
• 資料加解密與雜湊運算
• 算術邏輯轉換
• 攤平複雜的 JSON 結構
• 文章 / 文本分詞(Tokenization)
• 陣列與列表展開(類似 explode()

接下來,我們透過幾個實際範例,說明怎麼編寫 UDF/UDTF。

UDF

提示

useArrow=True 這個參數,可以省下過程中資料格式轉換所耗費的時間。沒有加上這個參數的 UDF,預設會依照以下順序執行:JVM 把 Java 物件轉成 Pickle 格式,丟給 Python 端解開,算完之後再 Pickle 回 JVM。

加上這個參數之後,Apache Arrow 會在記憶體中採用統一的欄位化(Columnar)格式,讓 Java 與 Python 可以直接理解同一塊記憶體,幾乎省去了格式轉換所需的時間。

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
from pyspark.sql import SparkSession
from pyspark.sql.functions import udf, col
from pyspark.sql.types import StructType, StructField, IntegerType, StringType, FloatType, DoubleType

# 定義 function
@udf(returnType=DoubleType(), useArrow=True)
def channel_bonus(revenue: int, channel: str):
    if channel == "POS":
        return revenue * 1.1
    elif channel == "Online":
        return revenue * 1.5
    else:
        return revenue

# Schema 設定
schema = StructType([
    StructField("Months", IntegerType(), nullable=False),
    StructField("Revenue", IntegerType(), nullable=False),
    StructField("Channel", StringType(), nullable=False)
])

# Spark 會話建立
spark = SparkSession.builder.appName("Spark_udf").getOrCreate()

# 資料集
data = [
    (1, 100, "POS"),
    (1, 150, "Online"),
    (2, 200, "POS"),
    (2, 250, "App"),
    (3, 300, "Online"),
    (3, 350, "Wholesale"),
    (1, 400, "App"),
    (2, 450, "POS")
]

# 轉換至 Spark DataFrame
ps_df = spark.createDataFrame(data, schema)
ps_df.show()

# 套用函數
ps_df.withColumn('bonus', channel_bonus(col('Revenue'), col('Channel'))).show()

udf

從結果可以看到,在定義 function 時需要先明確定義回傳的型別。定義完成後,就能在 DataFrame 操作中直接透過 function_name(p1, p2) 的方式傳入參數使用。

在 Spark SQL 中使用 UDF

在 Spark SQL 查詢中使用自訂函數時,有兩件事要特別注意:

  1. 把 DataFrame 註冊為 temp view
  2. 把 function 註冊給 SQL 查詢使用
 1
 2
 3
 4
 5
 6
 7
 8
 9
10
# 若要在 Spark SQL 查詢中使用

# 建立暫存表
ps_df.createOrReplaceTempView("q1_revenue")

# 註冊 function
spark.udf.register("channel_bonus", channel_bonus)

# SQL 查詢
spark.sql("select Months, Revenue, Channel, channel_bonus(Revenue, Channel) as bonus from q1_revenue").show()

udf_sql

UDTF

UDTF 的使用情境,比較偏向「把一個巢狀結構的資料,展開成多列結果」這類需求,光看文字可能有點抽象,我們直接看範例。

下面的程式碼一樣需要先定義好回傳格式,差別在於這裡要定義一個 class,並實作一個名稱固定為 eval 的方法,結尾則是用 yield 取代 return。因為執行時需要保留原本的欄位,所以我們透過 Spark SQL 查詢來呼叫;又因為需要對每一列資料逐一展開,所以搭配 LATERAL 讓它可以逐列取值。

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
from pyspark.sql import SparkSession
from pyspark.sql.functions import udf, col, lit
from pyspark.sql.types import StructType, StructField, IntegerType, StringType, FloatType, DoubleType
from pyspark.sql.functions import udtf
import re

# 定義函數
@udtf(returnType="keyword: string, length: int")
class ExtractHashtags:
    def eval(self, sentence: str):
        keywords = re.findall(r"#([\w-]+)", sentence)
        for keyword in keywords:
            yield (keyword, len(keyword))


# Spark 會話建立
spark = SparkSession.builder.appName("Spark_udtf").getOrCreate()

# Schema 設定
schema = StructType([
    StructField("sentence", StringType(), nullable=False)
])

data = [
    ("Hello world this is #big-data channel",),
    ("How can we use #Spark to process data",),
    ("Maybe we need to use some tool such as #airflow #redshift",)
]

ps_df = spark.createDataFrame(data, schema)

spark.udtf.register("extract_hashtags", ExtractHashtags)
ps_df.createOrReplaceTempView("ps_df")

spark.sql("""
    SELECT *
    FROM ps_df, LATERAL extract_hashtags(sentence)
""").show()

udtf

結語

這些工具都是自訂處理邏輯的方式,其中有些概念可能需要多練幾次才能真正掌握,很建議自己動手操作幾遍來熟悉。也可以參考官網這篇文章,會有更完整的說明:UDF and UDTF 官方教學

目錄