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


<!--more-->

## 前言

{{< admonition info >}}
本篇文章接續前篇 **[Spark教學(三) - Spark DataFrame 轉換 Pandas DataFrame](https://as183789043.github.io/spark-3/)**（連結請依實際路徑調整），繼續往下編寫處理邏輯。如果手上還沒有可用的環境，建議先依照前面幾篇的步驟完成安裝，跟著本篇操作會更順暢。
{{< /admonition >}}

</br>

## 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)** 最核心的差異，在於「輸入對輸出的映射關係」與「回傳的資料結構」：

* **UDF** 是 **1 對 1**：傳入一列資料，吐出一個欄位值（純量 / Scalar Value）。
* **UDTF** 是 **1 對多**：傳入一列資料，透過 `yield` 展開成零到多列、多欄位的完整表格（Table）。

{{< admonition note >}}
值得一提的是，Python UDTF 是 **Spark 3.5** 才正式加入的新功能，跟很早就存在的 UDF 並不是同期的產物。如果你的叢集版本較舊，可能得先確認一下是否支援。
{{< /admonition >}}

### 📊 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` 查詢 |
| **典型應用場景** | • 字串轉換與清理（如：去空白、轉大寫）<br>• 資料加解密與雜湊運算<br>• 算術邏輯轉換 | • 攤平複雜的 JSON 結構<br>• 文章 / 文本分詞（Tokenization）<br>• 陣列與列表展開（類似 `explode()`） |

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

## UDF

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

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

```python
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](udf.png)

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

## 在 Spark SQL 中使用 UDF

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

1. 把 DataFrame 註冊為 temp view
2. 把 function 註冊給 SQL 查詢使用

```python
# 若要在 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](udf_sql.png)

## UDTF

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

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

```python
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](udtf.png)

## 結語

這些工具都是自訂處理邏輯的方式，其中有些概念可能需要多練幾次才能真正掌握，很建議自己動手操作幾遍來熟悉。也可以參考官網這篇文章，會有更完整的說明：[UDF and UDTF 官方教學](https://spark.apache.org/docs/latest/api/python/user_guide/udfandudtf.html)

---

> 作者: Rick  
> URL: https://as183789043.github.io/zh-tw/theme-document-spark-query-utf-utdf/  

