Menu

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.

  • Intermediate
  • 18 min read
  • Updated Oct 2026
On this page
  1. Sample data
  2. UDFs and pandas UDFs
  3. What a UDF is and why it costs more
  4. Arrow-optimised Python UDFs in Spark 4.2
  5. NULLs and exceptions
  6. Using UDFs from SQL
  7. pandas UDFs (vectorised UDFs)
  8. Function APIs: applyInPandas and mapInPandas
  9. Table functions (UDTFs)
  10. Choosing
  11. Pitfalls
  12. In interviews
  13. The pandas API on Spark
  14. What it is
  15. How it works: the index
  16. Moving between APIs
  17. Pitfalls
  18. When to use it
  19. In interviews
  20. Practice questions
  21. 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:

  1. serialise batches of rows from the JVM and send them to a separate Python worker process on each executor;
  2. run your function in Python, one row at a time for a plain UDF;
  3. 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:

  1. Built-in functions (pyspark.sql.functions): strings, dates, regex, JSON, arrays and maps (transform, filter, aggregate), when/otherwise.
  2. pandas UDFs (or Arrow UDFs) when you need Python: vectorised over batches.
  3. 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.functions first; 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 mapInPandas with 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 a TypeError in 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_frames is True), 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.
  • applyInPandas needs each group in memory and mapInPandas and 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.

By DataDank Editorial · Last reviewed Oct 2026 · All examples run on PySpark 4.2.0 in local mode with pandas 3.0.6 and PyArrow 25.0.1 (PySpark 4.2 prints a FutureWarning that pandas 3 is not yet fully supported; warnings are suppressed in the examples). Arrow serialisation for regular Python UDFs is on by default from Spark 4.2; in earlier versions you opt in with useArrow=True or spark.sql.execution.pythonUDF.arrow.enabled.

Progress is saved in this browser only. No account needed.

Search
Filter by type