PySpark courseLesson 8 of 10
PySpark course · Lesson 8 of 10
PySpark UDFs, pandas UDFs and the Pandas API on Spark
When to write Python UDFs in PySpark, how Arrow and pandas UDFs reduce their cost, how applyInPandas and UDTFs work, and when the pandas API on Spark fits.
On this page
- Sample data
- UDFs and pandas UDFs
- What a UDF is and why it costs more
- Arrow-optimised Python UDFs in Spark 4.2
- NULLs and exceptions
- Using UDFs from SQL
- pandas UDFs (vectorised UDFs)
- Function APIs: applyInPandas and mapInPandas
- Table functions (UDTFs)
- Choosing
- Pitfalls
- In interviews
- The pandas API on Spark
- What it is
- How it works: the index
- Moving between APIs
- Pitfalls
- When to use it
- In interviews
- Practice questions
- Key takeaways
Spark’s built-in functions cover most transformations, but sometimes you need Python: a parsing library, a scoring model, business logic that is clearer in code. PySpark offers several ways to run Python on a DataFrame, and they differ a lot in speed and safety. This lesson covers plain UDFs, pandas (vectorised) UDFs and their grouped and map variants, table functions, and the pandas API on Spark, which lets pandas-style code run on a cluster.
Sample data
The setup also points Spark’s Python workers at the same interpreter as the driver. Mismatched Python environments between driver and workers are a classic cause of “module not found” errors inside UDFs.
import os, sys, re, io, contextlib, warnings
os.environ.setdefault("PYSPARK_PYTHON", sys.executable) # workers use the driver's Python
warnings.filterwarnings("ignore")
import pandas as pd
from pyspark.sql import SparkSession, functions as F
from pyspark.sql.types import StringType, IntegerType
spark = (SparkSession.builder.master("local[2]").appName("udfs")
.config("spark.sql.shuffle.partitions", "4").getOrCreate())
spark.sparkContext.setLogLevel("ERROR")
def plan(df):
buf = io.StringIO()
with contextlib.redirect_stdout(buf):
df.explain()
text = re.sub(r"#\d+L?", "", buf.getvalue())
print(re.sub(r", \[plan_id=\d+\]", "", text).strip())
orders = spark.createDataFrame(
[(1, "asha", "IN-560001", 120.0), (2, "ben", "UK-SW1A", 15.0),
(3, None, "IN-110001", 35.5), (4, "dara", None, 210.0)],
"order_id INT, customer STRING, postcode STRING, amount DOUBLE")
UDFs and pandas UDFs
What a UDF is and why it costs more
A user-defined function (UDF) wraps a Python function so Spark can apply it to a column. Spark’s engine runs in the JVM, so to call Python it must:
- serialise batches of rows from the JVM and send them to a separate Python worker process on each executor;
- run your function in Python, one row at a time for a plain UDF;
- serialise the results back to the JVM.
On top of that cost, the optimiser cannot see inside the function. It cannot push a filter on the result down to the file scan, cannot reorder around it, and cannot generate code for it.
@F.udf(returnType=StringType())
def country_of(postcode):
if postcode is None:
return None
return postcode.split("-")[0]
orders.select("order_id", country_of("postcode").alias("country")).show()
plan(orders.select(country_of("postcode").alias("country")))
plan(orders.select(F.split_part("postcode", F.lit("-"), F.lit(1)).alias("country")))
+--------+-------+
|order_id|country|
+--------+-------+
| 1| IN|
| 2| UK|
| 3| IN|
| 4| NULL|
+--------+-------+
== Physical Plan ==
*(2) Project [pythonUDF0 AS country]
+- ArrowEvalPython [country_of(postcode)], [pythonUDF0], 101
+- *(1) Project [postcode]
+- *(1) Scan ExistingRDD[order_id,customer,postcode,amount]
== Physical Plan ==
*(1) Project [element_at(stringsplitsql(postcode, -), 1, Some(), false) AS country]
+- *(1) Scan ExistingRDD[order_id,customer,postcode,amount]
The UDF adds a separate ArrowEvalPython operator, which breaks whole-stage code generation into two stages. The built-in split_part is a single JVM projection. The rule of thumb in order of preference:
- Built-in functions (
pyspark.sql.functions): strings, dates, regex, JSON, arrays and maps (transform,filter,aggregate),when/otherwise. - pandas UDFs (or Arrow UDFs) when you need Python: vectorised over batches.
- Plain Python UDFs for row-by-row logic that cannot be vectorised.
The optimiser cannot look through the UDF even for a simple filter on its output. Here the UDF is evaluated twice, once for the filter and once for the projection:
plan(orders.select(country_of("postcode").alias("c")).filter("c = 'IN'"))
== Physical Plan ==
*(3) Project [pythonUDF0 AS c]
+- ArrowEvalPython [country_of(postcode)], [pythonUDF0], 101
+- *(2) Project [postcode]
+- *(2) Filter (pythonUDF0 = IN)
+- ArrowEvalPython [country_of(postcode)], [pythonUDF0], 101
+- *(1) Project [postcode]
+- *(1) Scan ExistingRDD[order_id,customer,postcode,amount]
UDFs are assumed to be deterministic. If yours is not (random numbers, current time, calls to a service), mark it with .asNondeterministic() so Spark does not duplicate or reorder calls in ways that change results.
Arrow-optimised Python UDFs in Spark 4.2
Plain UDFs originally exchanged data with Python using pickled rows. Since Spark 4.2, Arrow (a columnar in-memory format) is used by default for this transfer (spark.sql.execution.pythonUDF.arrow.enabled = true; earlier versions needed useArrow=True). Your function still runs once per row with ordinary Python values; only the transfer is faster. The setting is read when a UDF is created, so a UDF defined while it is off shows the older operator:
print(spark.conf.get("spark.sql.execution.pythonUDF.arrow.enabled"))
spark.conf.set("spark.sql.execution.pythonUDF.arrow.enabled", "false")
country_pickled = F.udf(lambda p: p.split("-")[0] if p else None, StringType()) # created while off
plan(orders.select(country_pickled("postcode")))
spark.conf.set("spark.sql.execution.pythonUDF.arrow.enabled", "true")
true
== Physical Plan ==
*(2) Project [pythonUDF0 AS <lambda>(postcode)]
+- BatchEvalPython [<lambda>(postcode)], [pythonUDF0]
+- *(1) Project [postcode]
+- *(1) Scan ExistingRDD[order_id,customer,postcode,amount]
Arrow serialisation also changes type coercion, which matters when the declared return type does not match what the function returns. With pickled transfer, a mismatch quietly becomes NULL; with Arrow, values are converted (here, doubles truncated to integers):
legacy = F.udf(lambda x: x * 2, IntegerType(), useArrow=False)
arrow = F.udf(lambda x: x * 2, IntegerType(), useArrow=True)
orders.select("amount", legacy("amount").alias("legacy"), arrow("amount").alias("arrow")).show()
+------+------+-----+
|amount|legacy|arrow|
+------+------+-----+
| 120.0| NULL| 240|
| 15.0| NULL| 30|
| 35.5| NULL| 71|
| 210.0| NULL| 420|
+------+------+-----+
Neither result is what you want: 35.5 * 2 = 71.0 was silently truncated in one case and lost in the other. Declare the return type that matches what the function returns (DoubleType here), and test it. When upgrading to Spark 4.2, UDFs that relied on mismatches producing NULL can change behaviour.
NULLs and exceptions
Spark passes NULL as None. A function that does not handle it raises, and one bad row fails the whole job:
upper_unsafe = F.udf(lambda s: s.upper(), StringType())
try:
orders.select(upper_unsafe("customer")).collect()
except Exception as e:
print(type(e).__name__, "AttributeError" in str(e))
PythonException True
Handle None explicitly, decide what to return for invalid input, and unit-test the plain Python function before wrapping it.
Using UDFs from SQL
spark.udf.register(name, f) makes a UDF callable in SQL:
spark.udf.register("country_of", country_of)
orders.createOrReplaceTempView("orders")
spark.sql("SELECT order_id, country_of(postcode) AS country FROM orders WHERE order_id < 3").show()
+--------+-------+
|order_id|country|
+--------+-------+
| 1| IN|
| 2| UK|
+--------+-------+
pandas UDFs (vectorised UDFs)
A pandas UDF receives a whole batch of values as a pandas.Series (transferred with Arrow) and returns a Series of the same length. The per-row Python overhead disappears and you can use vectorised pandas and NumPy operations. The kind of pandas UDF is inferred from the type hints:
| Type hints | Kind | Use |
|---|---|---|
pd.Series -> pd.Series |
Scalar | Column-to-column transformation |
Iterator[pd.Series] -> Iterator[pd.Series] |
Scalar iterator | Same, with expensive setup done once per partition |
pd.Series -> scalar |
Grouped aggregate | Custom aggregation in groupBy().agg() or windows |
@F.pandas_udf("double")
def with_tax(amount: pd.Series) -> pd.Series:
return amount * 1.18
@F.pandas_udf("string")
def country_vec(postcode: pd.Series) -> pd.Series:
return postcode.str.split("-").str[0]
orders.select("order_id", F.round(with_tax("amount"), 2).alias("with_tax"),
country_vec("postcode").alias("country")).show()
plan(orders.select(with_tax("amount")))
+--------+--------+-------+
|order_id|with_tax|country|
+--------+--------+-------+
| 1| 141.6| IN|
| 2| 17.7| UK|
| 3| 41.89| IN|
| 4| 247.8| NULL|
+--------+--------+-------+
== Physical Plan ==
*(2) Project [pythonUDF0 AS with_tax(amount)]
+- ArrowEvalPython [with_tax(amount)], [pythonUDF0], 200
+- *(1) Project [amount]
+- *(1) Scan ExistingRDD[order_id,customer,postcode,amount]
The iterator form lets you load something expensive, such as a model or a lookup table, once per partition rather than once per batch:
from typing import Iterator
@F.pandas_udf("double")
def scaled(batches: Iterator[pd.Series]) -> Iterator[pd.Series]:
factor = 0.5 # expensive setup (a model, a lookup table) runs once per partition
for amount in batches:
yield amount * factor
orders.select("order_id", scaled("amount").alias("half")).show()
+--------+-----+
|order_id| half|
+--------+-----+
| 1| 60.0|
| 2| 7.5|
| 3|17.75|
| 4|105.0|
+--------+-----+
A grouped-aggregate pandas UDF reduces each group’s Series to one value:
@F.pandas_udf("double")
def median_amount(amount: pd.Series) -> float:
return float(amount.median())
with_country = orders.withColumn("country", F.split_part("postcode", F.lit("-"), F.lit(1)))
(with_country.groupBy("country").agg(median_amount("amount").alias("median"))
.orderBy("country").show())
+-------+------+
|country|median|
+-------+------+
| NULL| 210.0|
| IN| 77.75|
| UK| 15.0|
+-------+------+
Built-in percentile_approx or median would be cheaper here; a grouped-aggregate UDF needs all values of a group in memory on one worker.
Function APIs: applyInPandas and mapInPandas
These take a regular Python function over pandas DataFrames, not a column expression:
groupBy(...).applyInPandas(func, schema): each group arrives as one pandas DataFrame; return any number of rows. Good for per-group models or logic that needs the whole group.mapInPandas(func, schema): an iterator of pandas DataFrames per partition; return any number of rows. Good for row-wise enrichment that adds or removes rows.
def top_order(pdf: pd.DataFrame) -> pd.DataFrame:
return pdf.nlargest(1, "amount")[["country", "order_id", "amount"]]
(with_country.groupBy("country")
.applyInPandas(top_order, schema="country STRING, order_id INT, amount DOUBLE")
.orderBy("country").show())
def add_flag(batches):
for pdf in batches:
pdf["big"] = pdf["amount"] > 100
yield pdf
orders.mapInPandas(add_flag,
schema="order_id INT, customer STRING, postcode STRING, amount DOUBLE, big BOOLEAN").show()
+-------+--------+------+
|country|order_id|amount|
+-------+--------+------+
| NULL| 4| 210.0|
| IN| 1| 120.0|
| UK| 2| 15.0|
+-------+--------+------+
+--------+--------+---------+------+-----+
|order_id|customer| postcode|amount| big|
+--------+--------+---------+------+-----+
| 1| asha|IN-560001| 120.0| true|
| 2| ben| UK-SW1A| 15.0|false|
| 3| NULL|IN-110001| 35.5|false|
| 4| dara| NULL| 210.0| true|
+--------+--------+---------+------+-----+
applyInPandas shuffles by the grouping key and loads each whole group into memory on one worker; a skewed key (one country with most orders) can exhaust it. (The top-order problem itself is better solved with a window and row_number.)
Table functions (UDTFs)
A Python UDTF (Spark 3.5+) returns a table instead of a single value: its eval method yields zero or more rows per input. Use one when one input row should produce several output rows, for example exploding a custom format:
from pyspark.sql.functions import udtf
@udtf(returnType="order_id INT, part STRING")
class SplitPostcode:
def eval(self, order_id: int, postcode: str):
if postcode:
for part in postcode.split("-"):
yield order_id, part
SplitPostcode(F.lit(1), F.lit("IN-560001")).show()
+--------+------+
|order_id| part|
+--------+------+
| 1| IN|
| 1|560001|
+--------+------+
Registered with spark.udtf.register, a UDTF can be called in SQL with LATERAL to apply it to each row of a table. For simple splitting, built-ins such as explode(split(...)) remain faster (see nested data and explode). Recent releases also add Arrow UDFs (arrow_udf, arrow_udtf) that work directly on PyArrow arrays instead of pandas objects.
Choosing
| Option | Runs | Relative cost | When |
|---|---|---|---|
| Built-in functions | JVM, code-generated | Lowest | Always first |
| pandas UDF / Arrow UDF | Python, vectorised batches | Moderate | Python logic that can be vectorised |
| Python UDF (Arrow transfer) | Python, row by row | Higher | Row-wise logic that cannot be vectorised |
applyInPandas |
Python, whole group | Depends on group size | Per-group models or logic needing the whole group |
mapInPandas / UDTF |
Python, per partition / per row | Moderate | Changing the number of rows |
Pitfalls
- Writing a UDF for something a built-in does. Check
pyspark.sql.functionsfirst; it has hundreds of functions. - Wrong return type: values silently become NULL (pickled) or are coerced (Arrow).
- Not handling
None: one NULL fails the job. - Calling external services row by row from a UDF: thousands of concurrent requests, no batching, retries rerun calls. Use
mapInPandaswith batching, or do the enrichment outside Spark. - Environment mismatches: pandas, PyArrow and your libraries must be installed on every executor with compatible versions.
- Large closures: objects referenced by the function are serialised with it to every task; broadcast large lookups instead.
In interviews
“Why avoid Python UDFs?” is one of the most common PySpark questions. A strong answer: serialisation between JVM and Python, row-by-row execution, opacity to Catalyst (no pushdown, no code generation), and NULL/type pitfalls; then the alternatives in order (built-ins, pandas UDFs, Arrow) and when a UDF is still reasonable. Mention that Spark 4.2 uses Arrow for regular UDF transfer by default, which reduces but does not remove the cost. See the interview answer.
The pandas API on Spark
What it is
pyspark.pandas (formerly the Koalas project, part of PySpark since 3.2) implements much of the pandas API on top of Spark DataFrames. Code that looks like pandas runs distributed, so a team with pandas notebooks can scale them without rewriting everything in the DataFrame API.
import pyspark.pandas as ps
psdf = ps.DataFrame({"region": ["north", "north", "south", "south", "north"],
"rep": ["asha", "ben", "chen", "dara", "asha"],
"amount": [100, 150, 200, 50, 150]})
print(type(psdf).__name__)
print(psdf.groupby("region")["amount"].sum().sort_index())
print(psdf[psdf.amount > 100].sort_values("amount").head(3))
DataFrame
region
north 400
south 250
Name: amount, dtype: int64
region rep amount
1 north ben 150
4 north asha 150
2 south chen 200
How it works: the index
pandas DataFrames always have an index, and Spark DataFrames have no row order or index. pandas-on-Spark therefore keeps an extra index column and translates each operation into Spark plans. Operations that align on the index, such as assigning a column computed by groupby().transform, become joins on that index:
psdf["share"] = psdf.groupby("region")["amount"].transform(lambda s: s / s.sum()).round(2)
print(psdf.sort_values(["region", "rep", "amount"]))
buf = io.StringIO()
with contextlib.redirect_stdout(buf):
psdf.spark.explain()
print("plan contains a join:", "Join" in buf.getvalue())
region rep amount share
0 north asha 100 0.25
4 north asha 150 0.38
1 north ben 150 0.38
2 south chen 200 0.80
3 south dara 50 0.20
plan contains a join: True
When a Spark DataFrame is converted with pandas_api(), Spark must invent an index. The default type is distributed-sequence: a globally sequential 0, 1, 2… index computed in a distributed way, which still costs an extra pass over the data. distributed is cheaper (unique, not sequential); sequence uses a window without partitioning and puts all data in one partition, so avoid it on large data. Better still, name an existing column as the index (pandas_api(index_col="order_id")).
print(ps.get_option("compute.default_index_type"))
sdf = psdf.to_spark()
sdf.printSchema()
back = orders.pandas_api(index_col="order_id")
print(back.dtypes)
print(back.index)
distributed-sequence
root
|-- region: string (nullable = false)
|-- rep: string (nullable = false)
|-- amount: long (nullable = false)
|-- share: double (nullable = true)
customer str
postcode str
amount float64
dtype: object
Index([1, 2, 3, 4], dtype='int32', name='order_id')
to_spark() drops the pandas-on-Spark index (the schema above has no index column) unless you pass index_col to keep it as a column. With index_col="order_id", the column becomes the index and no default index is generated.
Moving between APIs
| Call | Moves data to the driver? |
|---|---|
spark_df.pandas_api() / psdf.to_spark() |
No, both stay distributed |
spark_df.toPandas() / psdf.to_pandas() |
Yes: everything is collected into one pandas DataFrame |
ps.from_pandas(pdf) / spark.createDataFrame(pdf) |
Data starts on the driver and is distributed |
toPandas() uses Arrow by default in Spark 4 (spark.sql.execution.arrow.pyspark.enabled), which makes the transfer faster but does not reduce the memory needed on the driver.
Pitfalls
- Not all of pandas is supported, and some functions behave differently; for example, the string shortcut
transform("sum")used above in pandas raises aTypeErrorin pandas-on-Spark, so a function was passed instead. Check the API reference for each function you rely on. - Order-dependent operations (
head,iloc,shift, cumulative functions) need an ordering and can trigger expensive sorts or single-partition windows. - Combining two DataFrames with different origins (
psdf1.a + psdf2.b) requires a join on the index; Spark 4 allows it by default (compute.ops_on_diff_framesisTrue), but it is a hidden join. - Version support: PySpark releases support a range of pandas versions; with a newer pandas than tested (PySpark 4.2 warns about pandas 3), some behaviour may differ.
When to use it
| Use the pandas API on Spark | Use the DataFrame API |
|---|---|
| Scaling an existing pandas notebook or library with minimal rewrite | New production pipelines |
| Exploratory work by pandas users | Performance-critical jobs where you want to see and control the plan |
| Using pandas-style plotting and statistics on large data | Streaming (pandas-on-Spark is batch only) |
In interviews
Questions are usually “how would you scale pandas code?” or “what is the difference between toPandas() and the pandas API on Spark?”. Strong answers explain that pandas-on-Spark stays distributed while toPandas collects to the driver, mention the default index cost and its options, and say that new pipelines are usually written with the DataFrame API.
Practice questions
Why is a Python UDF slower than an equivalent built-in function?
Data must be serialised from the JVM to a Python worker and back, the function runs row by row in the Python interpreter, and the optimiser cannot see inside it, so it cannot push filters through it, reorder it, or include it in whole-stage code generation. Arrow transfer (default for UDFs in Spark 4.2) reduces the serialisation cost only.
A UDF declared as IntegerType returns float values. What happens?
With the classic pickled transfer, values that do not match the declared type become NULL silently. With Arrow-optimised UDFs (default in Spark 4.2), they are coerced, for example truncated to integers. Either way the data is wrong; declare the type the function actually returns and test it.
When would you choose a pandas UDF over a Python UDF, and when applyInPandas?
A pandas UDF processes Arrow batches as pandas Series, so vectorisable logic runs much faster than row-by-row Python. Use applyInPandas when the logic needs an entire group at once (for example fitting a model per customer), accepting that each group must fit in a worker’s memory and that the data is shuffled by the key.
How do you load a machine learning model once per executor task rather than once per row?
Use an iterator-of-Series pandas UDF (or mapInPandas) and load the model before looping over the batches, so it happens once per partition. Broadcasting the model file or caching it in a module-level variable in the worker are common complements.
What is the difference between toPandas() and pandas_api()?
toPandas() collects all rows into a single pandas DataFrame on the driver, which fails for data larger than driver memory. pandas_api() returns a pandas-on-Spark DataFrame that stays distributed and translates pandas-style calls into Spark plans.
Why can converting a Spark DataFrame with pandas_api() be slow, and how do you avoid it?
Spark DataFrames have no index, so pandas-on-Spark attaches a default one. The default distributed-sequence index needs an extra distributed computation, and sequence forces data into one partition. Specify an existing unique column with index_col, or use the distributed index type when sequential values are not needed.
Key takeaways
- Prefer built-in functions; Python UDFs add serialisation, row-by-row execution and an optimiser blind spot.
- Spark 4.2 uses Arrow for regular UDF transfer by default, which is faster but changes how return-type mismatches behave.
- Handle
None, declare correct return types, mark non-deterministic UDFs, and test the plain Python function. - pandas UDFs are vectorised; their kind comes from type hints, and the iterator form amortises setup per partition.
applyInPandasneeds each group in memory andmapInPandasand UDTFs change the number of rows; use them when column functions cannot.- The pandas API on Spark stays distributed but maintains an index that can cost extra passes;
toPandas()collects everything to the driver.
Progress is saved in this browser only. No account needed.

