AWS Training
Modules Listen Certification
0:00 0:00

← PySpark

Starts this lesson and continues through 9 more to the end of the course.

The Python/JVM boundary

The sentence that explains everything else

Your Python code mostly does not run on your data.

When you write df.filter(col("x") > 5), no filtering happens. You have appended a node to a logical plan. That plan is handed to the JVM, optimised there, and executed there — across executors that never run a line of your Python.

Which raises the obvious question: when does Python touch data? The answer is a short list, and every item on it is expensive:

Crossing What moves
A Python UDF Every row, serialised out to a Python worker and back
collect() All result rows, to the driver's Python process
toPandas() All result rows, to the driver, materialised as pandas
take(n) / show() n rows to the driver — bounded, so usually fine
foreach / mapPartitions (RDD API) Data to Python workers

Everything else you write is plan construction. That's why PySpark can feel instant right up until the moment it doesn't.

Why Arrow exists

Moving data between the JVM and Python is the core problem, and Apache Arrow is the answer to it:

"Apache Arrow is an in-memory columnar data format that is used in Spark to efficiently transfer data between JVM and Python processes." — Apache Arrow in PySpark (verified 2026-08-12)

Note what that sentence concedes: transfer between JVM and Python is expensive enough that a dedicated columnar format exists purely to make it less bad. Arrow narrows the boundary cost. It does not remove the boundary.

Lesson 3 covers the configuration. What matters here is the shape: the boundary is the thing, and Arrow is the mitigation.

What "the driver" is, and how you kill it

The driver is a single process running your Python script. Executors are many processes doing the work. Your SparkSession lives in the driver.

Three ways to kill it, in descending order of frequency:

# 1. collect() on something large — pulls every row into one process
rows = df.collect()                     # 200 GB DataFrame → driver OOM

# 2. toPandas() on something large — same, plus a pandas materialisation
pdf = df.toPandas()

# 3. Building a huge object in the driver and closing over it
big_dict = load_500mb_lookup()          # serialised to every task
df.withColumn("x", udf(lambda v: big_dict.get(v))(col("v")))

⚠️ The failure mode is not a nice error. It's the driver going unresponsive, then dying, and the job failing with something that looks unrelated. On Glue in particular you'll see the job fail with little explanation, because the process that would have reported the problem is the one that died.

The rule: nothing unbounded ever comes back to the driver.

df.limit(1000).toPandas()               # bounded — fine
df.write.parquet(path)                  # distributed — fine
df.count()                              # one number — fine
df.collect()                            # unbounded — no

If you need all the data somewhere, write it out and read it from there. The driver is a coordinator, not a data store.

withColumn in a loop

This one looks innocent and is a genuine production problem:

# BAD: 200 iterations → a plan 200 nodes deep
for c in columns:
    df = df.withColumn(c, trim(col(c)))

Each withColumn returns a new DataFrame with another projection node. Two hundred of those produce a deeply nested logical plan, and plan analysis and optimisation are not free — they happen in the driver, single-threaded, before any work starts. Teams have seen jobs where planning took longer than execution.

# GOOD: one projection, one node
df = df.select([trim(col(c)).alias(c) if c in columns else col(c) for c in df.columns])

The general principle: build the expression, then apply it once. Prefer a single select with a list of expressions over repeated withColumn. Same for chained withColumnRenamed.

⚠️ I have not verified a specific documented threshold at which plan depth becomes a problem — the number will depend on your Spark version and the complexity of each expression. Treat "loops that call withColumn" as a smell to be refactored, and measure if you want a threshold for your environment.

The lookup table question

A common shape: you have a small reference table and want to enrich a large one.

# Option A: Python dict + UDF — crosses the boundary for every row
lookup = {"a": 1, "b": 2}
df.withColumn("v", udf(lambda k: lookup.get(k), IntegerType())(col("k")))

# Option B: broadcast join — stays entirely on the JVM
df.join(broadcast(lookup_df), on="k", how="left")

Option B is almost always right. The dict-plus-UDF version serialises the dict to every task, then round-trips every row through a Python process. The broadcast join ships the small table once and does the lookup in JVM code.

This is the single highest-value refactor in most PySpark codebases, and it's a standard interview question. The answer they want is the mechanism: the UDF forces a row-by-row crossing of the Python/JVM boundary; the broadcast join doesn't cross it at all.

Reading a script for boundary crossings

Train yourself to scan for them. In this script, where does Python touch data?

df = spark.read.parquet("s3://bucket/events/")        # plan
df = df.filter(col("event_date") >= "2026-01-01")     # plan
df = df.withColumn("clean", trim(col("raw")))         # plan — built-in, JVM
df = df.withColumn("score", my_udf(col("clean")))     # ⚠️ CROSSING, every row
agg = df.groupBy("user_id").agg(sum("score"))         # plan
top = agg.orderBy(desc("sum(score)")).limit(100)      # plan
rows = top.collect()                                  # ⚠️ CROSSING — but bounded to 100
top.write.parquet("s3://bucket/out/")                 # distributed write

Two crossings. One is bounded to 100 rows and harmless. The other runs per row over the whole dataset and is where the job's time goes.

That scan takes ten seconds and finds most PySpark performance problems.

Check yourself

  1. df.filter(...) — what has happened after that line runs?
  2. Name three operations that move data into the driver.
  3. Why is a Python dict plus a UDF worse than a broadcast join?
  4. What's wrong with calling withColumn in a loop over 200 columns?
  5. Why does Arrow exist?
Answers
  1. Nothing has been filtered. A node was appended to a logical plan. Execution happens later, on the JVM, when an action is called.
  2. collect(), toPandas(), and take(n)/show() (the last is bounded, so usually safe). Anything unbounded is a driver risk.
  3. The dict is serialised to every task, and the UDF round-trips every row through a Python worker — a per-row crossing of the JVM/Python boundary. A broadcast join ships the small table once and performs the lookup in JVM code, never crossing.
  4. Each call adds another projection node, producing a deeply nested plan that the driver must analyse and optimise single-threaded before any work starts. Build one select with a list of expressions instead.
  5. Because transferring data between JVM and Python processes is expensive. Arrow is "an in-memory columnar data format used in Spark to efficiently transfer data between JVM and Python processes" — it narrows the cost of the boundary, but doesn't remove the boundary.

Teaching this section

Next →Lazy evaluation, transformations, and actions