공부를 하다/Databricks

Day 09. Spark UDF + Higher Order Functions

Banaaan 2026. 8. 24. 12:49

pandas의 apply()나 SQL UDF를 써본 사람에게도 Spark UDF는 낯설다. 분산 환경이라는 특성 때문에 등록 방식과 성능 특성이 다르다.


Spark UDF

왜 pandas처럼 apply()를 못 쓰나?

Spark는 데이터가 여러 서버에 분산되어 있다. apply()는 단일 머신에서 순서대로 실행하는 방식이라 분산 환경에서 동작하지 않는다. UDF는 Spark가 각 파티션에 함수를 전달해서 분산 실행할 수 있도록 만든 것이다.


Python UDF 만들기

from pyspark.sql.functions import udf
from pyspark.sql.types import StringType

@udf(returnType=StringType())
def grade_label(score):
    if score >= 90:
        return "A"
    elif score >= 80:
        return "B"
    else:
        return "C"

df.withColumn("등급", grade_label("score"))

pandas apply()와 다른 점:
1. 반환 타입을 반드시 명시해야 한다. Spark는 분산 환경이라 타입을 미리 알아야 한다.
2. apply()가 아닌 withColumn() 안에서 사용한다.


Python UDF의 단점 — 시험 포인트

Spark (JVM) → 데이터를 Python으로 직렬화 → Python에서 함수 실행 → 다시 JVM으로 역직렬화

이 변환 비용이 크다. Spark 내장 함수로 해결 가능하면 UDF 쓰지 말고 내장 함수를 써야 한다. 내장 함수로 불가능한 복잡한 로직일 때만 UDF를 사용한다.


Pandas UDF (Vectorized UDF) — 빠른 UDF

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

@pandas_udf(StringType())
def grade_label(score: pd.Series) -> pd.Series:
    return score.apply(lambda x: "A" if x >= 90 else ("B" if x >= 80 else "C"))

df.withColumn("등급", grade_label("score"))

입력/출력이 pandas Series다. 행 단위가 아닌 배치(묶음) 단위로 처리해서 Python UDF보다 훨씬 빠르다.

  Python UDF Pandas UDF
입력 값 하나 pandas Series
속도 느림 (직렬화 비용) 빠름 (Arrow 기반)
코드 단순 pandas 문법 필요

Higher Order Functions

배열(Array) 컬럼을 다루는 Spark 고유 함수다. 중첩 JSON이나 배열 컬럼이 있는 데이터를 처리할 때 사용한다.

df = spark.createDataFrame([
    (1, [10, 20, 30, 5]),
    (2, [100, 3, 50])
], ["id", "scores"])

transform() — 각 요소에 함수 적용

Python의 map(), pandas의 .apply()와 같은 개념이다.

from pyspark.sql.functions import transform

df.withColumn("doubled", transform("scores", lambda x: x * 2))
# id=1 → [20, 40, 60, 10]
# id=2 → [200, 6, 100]

filter() — 조건에 맞는 요소만 남기기

from pyspark.sql.functions import filter

df.withColumn("high_scores", filter("scores", lambda x: x >= 10))
# id=1 → [10, 20, 30]   (5 제외)
# id=2 → [100, 50]      (3 제외)

DataFrame의 filter()(행 필터)와 이름이 같지만, 이건 배열 안의 요소 필터다.


exists() — 하나라도 조건 만족하면 True

from pyspark.sql.functions import exists

df.withColumn("has_perfect", exists("scores", lambda x: x >= 100))
# id=1 → False
# id=2 → True

aggregate() — 배열을 하나의 값으로 줄이기

Python의 reduce()와 같다.

from pyspark.sql.functions import aggregate, lit

df.withColumn("total", aggregate("scores", lit(0), lambda acc, x: acc + x))
# id=1 → 65
# id=2 → 153

lit(0) = 초기값


SQL 방식

SELECT id,
       TRANSFORM(scores, x -> x * 2) AS doubled,
       FILTER(scores, x -> x >= 10) AS high_scores,
       EXISTS(scores, x -> x >= 100) AS has_perfect
FROM scores_table

SQL에서는 x -> x * 2 형태로 쓴다 (Python은 lambda x: x * 2).


한 줄 정리

transform()   → 각 요소 변환  (map과 동일)
filter()      → 조건 만족 요소만 유지
exists()      → 하나라도 조건 만족하면 True
aggregate()   → 배열을 하나의 값으로 축약  (reduce와 동일)

✓ Day 1: Spark 기초
✓ Day 2: Delta Lake
✓ Day 3: Auto Loader
✓ Day 4: Structured Streaming
✓ Day 5: Delta Live Tables
✓ Day 6: Workflows
✓ Day 7: Architecture + Notebooks + Git Folders
✓ Day 8: DataFrameReader API + Views + JOIN + 집계 + Window Functions
✓ Day 9: Spark UDF + Higher Order Functions