AWS Training
Modules Listen Certification
0:00 0:00

← PySpark

PS0 Cheat sheet — PySpark

Verified against the current PySpark documentation on spark.apache.org, 2026-08-12. ⚠️ Defaults are version-dependent. Check spark.version in your Glue/EMR environment and read the docs for THAT version.

The boundary — the whole model

YOUR PYTHON (driver)          │        JVM EXECUTORS
df = spark.read...            │
df = df.filter(...)           │   plan construction — nothing runs
df = df.withColumn(...)       │
df.write.parquet(...)    ────▶│   NOW it executes
──────────────────────────────────────────────────────
CROSSINGS:  udf() per row · collect() · toPandas() · take()/show() (bounded)

Every performance question = how often and how much are we crossing that line?

"Apache Arrow is an in-memory columnar data format that is used in Spark to efficiently transfer data between JVM and Python processes."

Lazy evaluation

Transformations (free) Actions (execute everything)
select filter withColumn join groupBy().agg() write count collect show take toPandas

Narrow (no shuffle): filter select withColumn Wide (shuffle, stage boundary): groupBy join distinct orderBy repartition

explain() — what to look for

  1. Exchange = shuffle. Count them. 2. Join strategy (broadcast vs shuffle).
  2. PushedFilters / partition pruning. 4. Scan size.

UDFs — the hierarchy

1. Built-in (pyspark.sql.functions)   ← exhaust this first
2. SQL expression / higher-order
3. pandas UDF (Arrow batches)
4. Plain Python UDF (row by row)
5. RDD mapPartitions (no optimiser)

Three costs of a plain Python UDF: per-row serialisation · a separate Python process outside the JVM heap · the optimiser goes blind (no predicate pushdown, no column pruning).

df.filter(is_valid_udf(col("region")))          # 🚨 no pushdown — reads everything
df.filter(col("region").isin("EU","US"))        # ✅ pushed to the scan

Dict + UDF → broadcast join. Highest-value refactor in most codebases.

Arrow config (documented values)

Key Documented
spark.sql.execution.arrow.pyspark.enabled pandas conversion; default true
spark.sql.execution.arrow.pyspark.fallback.enabled silent fallback to non-Arrow on error ⚠️
spark.sql.execution.arrow.maxRecordsPerBatch default 10,000 rows/batch
spark.sql.execution.arrow.pyspark.selfDestruct.enabled saves memory on toPandas
spark.sql.execution.pythonUDF.arrow.enabled Arrow for regular Python UDFs
spark.sql.session.timeZone defaults to JVM system local TZ

Types: all supported except ArrayType of TimestampType; MapType and ArrayType of nested StructType need PyArrow ≥ 2.0.0. Minimums: pandas 2.2.0, PyArrow 18.0.0.

pandas UDF types

Type Signature Constraint
Series → Series Series,… -> Series output length == input length
Iterator Iterator[Series] -> Iterator[Series] same total length; setup once per iterator
Iterator of tuples Iterator[Tuple[Series,…]] -> Iterator[Series] as above
Series → Scalar Series,… -> Any 🚨 "no partial aggregation, all data for a group or window loaded into memory"; unbounded window only

pandas Function APIs

API Note
groupby().applyInPandas() ⚠️ whole group into memory
mapInPandas() ✅ streams; arbitrary output length
groupby().cogroup().applyInPandas() ⚠️ whole cogroup into memory

Partitions and skew

Spark partitions = parallelism (1 task each). Storage partitions = S3 directories (partitionBy). Storage decides what you read; Spark decides how you process it.

# Diagnose
df.groupBy("key").count().orderBy(desc("count")).show(20)
df.groupBy(spark_partition_id().alias("pid")).count().orderBy(desc("count")).show(20)

Spark UI tell: stage max task duration ≫ median.

Causes: dominant key (NULL/0/-1/"UNKNOWN"/whale) · low-cardinality join key · unsplittable GZIP → one task · dominant group.

Fixes, in order: filter the junk key → broadcast() the small side → salt the hot key → repartition on a better key → AQE (⚠️ verify keys/defaults for your version).

Shuffle Can increase Note
repartition(n) Yes Yes keeps upstream parallelism
coalesce(n) No No 🚨 narrows upstream parallelism too — coalesce(1) can serialise the whole job

Small files: repartition before write; repartition("dt").write.partitionBy("dt").

Production

# transforms.py — pure, testable
def clean(df: DataFrame) -> DataFrame: ...
# job.py — thin I/O shell
def main(spark, in_path, out_path): ...
# tests: THIS LINE makes the suite fast
.config("spark.sql.shuffle.partitions", "2")

Assert in the job, not just in tests — tests use data you invented; assertions see today's data. Fail the job rather than write bad data.

Idempotency: mode("overwrite") ✅ · mode("append") 🚨 doubles on retry. ⚠️ overwrite + partitionBy semantics depend on a dynamic-overwrite setting — verify for your version; getting it wrong deletes data. Assume every job runs twice.

On AWS — do this first

print(spark.version)
print(spark.conf.get("spark.sql.execution.arrow.pyspark.enabled"))
print(spark.sparkContext.getConf().getAll())

Write the output into your runbook, then read the docs for that version.