Apache Spark courseLesson 4 of 5
Apache Spark course · Lesson 4 of 5
Spark Partitions, Shuffles and Data Skew
How Spark partitions data, what a shuffle does on disk and over the network, how to size shuffle partitions, and how to detect and fix skew with AQE and salting.
On this page
- Sample data
- Partitions and parallelism
- What it is
- Where partition counts come from
- How many partitions you want
- Pitfalls
- In interviews
- Shuffle internals
- What it is
- The map side: write sorted, partitioned files
- The reduce side: fetch and merge
- Who serves the files
- Pitfalls
- In interviews
- Understanding shuffle cost
- Where the time goes
- Ways to make shuffles cheaper
- Reading shuffle volume
- In interviews
- Repartition versus coalesce
- What they do
- The coalesce trap
- When to use which
- In interviews
- Choosing a partitioning strategy
- The three partitioners
- Two meanings of “partitioning”
- Choosing a key
- In interviews
- Tuning spark.sql.shuffle.partitions
- What it is
- With Adaptive Query Execution (the default since Spark 3.2)
- Sizing it by hand
- Pitfalls
- In interviews
- Detecting data skew
- What it is
- Detect it in the data
- Detect it in task metrics
- Common causes
- In interviews
- Salting for skew
- What it is
- Salting an aggregation (two-phase aggregation)
- Salting a join
- Pitfalls
- In interviews
- Skew join optimisation
- Options, in the order to try them
- Seeing AQE’s skew join
- Broadcast avoids the problem
- What AQE skew handling does not cover
- In interviews
- Practice questions
- Key takeaways
Most Spark performance problems come down to three connected ideas: how data is split into partitions, how much of it has to move in a shuffle, and whether the work is spread evenly or piled onto a few tasks by skew. This lesson works through each with measurements you can reproduce, then ends with the fixes interviewers expect you to know.
Sample data
A clickstream of 100,000 events over 5,000 users, deliberately skewed: 60% of events belong to user_id = 0 (think of a bot, a default value or a missing id). spark.local.dir points at a temporary folder so you can look at shuffle files later. api() reads the Spark UI REST API for task metrics.
import json, os, statistics, tempfile, urllib.request
from pyspark.sql import SparkSession, functions as F
local_dir = tempfile.mkdtemp(prefix="spark-local-")
spark = (SparkSession.builder.master("local[2]").appName("partitions-shuffles-skew")
.config("spark.local.dir", local_dir)
.config("spark.sql.shuffle.partitions", "8")
.getOrCreate())
sc = spark.sparkContext
_opener = urllib.request.build_opener(urllib.request.ProxyHandler({}))
def api(path):
port = sc.uiWebUrl.rsplit(":", 1)[1]
url = f"http://localhost:{port}/api/v1/applications/{sc.applicationId}/{path}"
return json.load(_opener.open(url))
def sizes(df):
"""Rows in each partition."""
return df.rdd.glom().map(len).collect()
events = (spark.range(0, 100_000, numPartitions=8)
.withColumn("user_id", F.when(F.col("id") % 10 < 6, F.lit(0)).otherwise(F.col("id") % 5000))
.withColumn("amount", (F.col("id") % 50).cast("int")))
users = (spark.range(0, 5000).withColumnRenamed("id", "user_id")
.withColumn("country", F.element_at(F.array(F.lit("uk"), F.lit("in"), F.lit("us")),
(F.col("user_id") % 3 + 1).cast("int"))))
Partitions and parallelism
What it is
A partition is a chunk of a dataset that one task processes. Spark runs one task per partition per stage, and each executor core runs one task at a time. So the partition count sets the maximum parallelism of a stage, and the core count sets how much of it you get at once.
print("partitions:", events.rdd.getNumPartitions())
print("rows per partition:", sizes(events))
print("task slots:", sc.defaultParallelism)
partitions: 8
rows per partition: [12500, 12500, 12500, 12500, 12500, 12500, 12500, 12500]
task slots: 2
Eight partitions on two slots run in four waves of two tasks.
Where partition counts come from
| Source | Partition count |
|---|---|
| File scan (Parquet, ORC, CSV, JSON) | Files are split into chunks of up to spark.sql.files.maxPartitionBytes (128 MB by default); small files are packed together, with each file charged an extra spark.sql.files.openCostInBytes |
spark.range, parallelize |
The number you pass, or spark.default.parallelism (total cores on most cluster managers) |
| After a DataFrame shuffle | spark.sql.shuffle.partitions (200 by default), then possibly coalesced by AQE |
repartition(n) / coalesce(n) |
n |
RDD shuffle (reduceByKey) |
The parent’s partition count unless you pass numPartitions |
The file rule is easy to see. Eight Parquet files of about 65 KB each are packed into two read partitions by default, and split into eight when the maximum is lowered to 64 KB:
path = os.path.join(tempfile.mkdtemp(), "events")
events.write.parquet(path)
print("files written:", len([f for f in os.listdir(path) if f.endswith(".parquet")]))
print("read partitions (128 MB max):", spark.read.parquet(path).rdd.getNumPartitions())
spark.conf.set("spark.sql.files.maxPartitionBytes", "64k")
print("read partitions (64 KB max):", spark.read.parquet(path).rdd.getNumPartitions())
spark.conf.unset("spark.sql.files.maxPartitionBytes")
files written: 8
read partitions (128 MB max): 2
read partitions (64 KB max): 8
How many partitions you want
Rules of thumb that hold up in practice:
- At least two to four times the total cores for big stages, so slow tasks do not leave cores idle and stragglers overlap with other work.
- Partitions of roughly 100–200 MB of input for scans, and shuffle partitions that each hold tens to a few hundred megabytes. Much smaller and scheduling overhead dominates (each task costs milliseconds to launch); much larger and tasks spill or run out of memory.
- Even sizes matter more than the count. One partition ten times bigger than the rest decides the stage’s duration.
Pitfalls
- Reading a folder of thousands of tiny files: the scan packs them into partitions, but listing and opening files dominates. Compact the files.
spark.rangeorparallelizewithout a partition count in tests: you getdefaultParallelism, which differs between laptop and cluster.- Checking
getNumPartitions()on a DataFrame before an action under AQE: the shuffle partition count may be coalesced at run time.
In interviews
“How do partitions relate to tasks and cores?” One task per partition per stage; one task per core at a time; so total cores cap concurrency and partition count caps parallelism. Mention the 128 MB file split default and the 200 shuffle partition default, and that AQE adjusts the latter.
Shuffle internals
What it is
A shuffle redistributes rows across partitions so that all rows with the same key end up in the same partition. It is how Spark implements groupBy, joins, distinct, repartition and window partitionBy. A shuffle splits the job into two stages: a map side that writes data and a reduce side that reads it.
The map side: write sorted, partitioned files
Each map task (one per input partition) computes a target partition for every row, by hashing the key modulo the number of shuffle partitions for hash partitioning. Spark’s sort-based shuffle then buffers rows in memory, sorts them by target partition id, spills sorted runs to local disk when the buffer fills, and merges them into one data file plus one index file per map task. The index file stores the byte offset where each reduce partition starts in the data file. Files go to spark.local.dir (or the cluster manager’s local directories), compressed with spark.io.compression.codec (lz4 by default).
You can see them. Four map tasks shuffling into eight partitions write four data files and four index files:
df = spark.range(0, 100_000, numPartitions=4).withColumn("k", F.col("id") % 100)
df.groupBy("k").count().collect()
shuffle_files = sorted((f, os.path.getsize(os.path.join(root, f)))
for root, _, files in os.walk(local_dir) for f in files
if f.startswith("shuffle_"))
for name, size in shuffle_files:
print(name, size)
shuffle_0_18_0.checksum.ADLER32 64
shuffle_0_18_0.data 1262
shuffle_0_18_0.index 72
shuffle_0_19_0.checksum.ADLER32 64
shuffle_0_19_0.data 1262
shuffle_0_19_0.index 72
shuffle_0_20_0.checksum.ADLER32 64
shuffle_0_20_0.data 1262
shuffle_0_20_0.index 72
shuffle_0_21_0.checksum.ADLER32 64
shuffle_0_21_0.data 1262
shuffle_0_21_0.index 72
The name is shuffle_<shuffleId>_<mapId>_<reduceId>. The map id is the map task’s unique attempt id (18 to 21 here, because earlier jobs in the session already used lower ids), and the reduce id is always 0 because one file holds all reduce partitions. Each index file is 72 bytes: nine 8-byte offsets marking the boundaries of eight reduce partitions. The checksum files let Spark diagnose corrupt shuffle blocks.
The reduce side: fetch and merge
Each reduce task asks the driver’s map output tracker where its partition’s blocks live, then fetches its slice from every map task’s file over the network (or from local disk if on the same executor), up to spark.reducer.maxSizeInFlight (48 MB by default) at a time. It then aggregates, joins or sorts the fetched rows, spilling to disk if they do not fit in execution memory.
Who serves the files
Shuffle files belong to the executor that wrote them. If that executor is removed, reducers cannot fetch its output and Spark must rerun those map tasks. Two mechanisms keep files available when executors come and go: an external shuffle service on each node (spark.shuffle.service.enabled, common on YARN), or shuffle tracking (spark.dynamicAllocation.shuffleTracking.enabled, on by default), which keeps executors holding live shuffle data alive. Both matter mostly with dynamic allocation, covered in the memory and tuning lesson.
Pitfalls
- Fetch failures (
FetchFailedException) usually mean an executor holding shuffle files died, often from memory pressure. Spark then reruns the parent stage, so the root cause is usually upstream of the error. - Local disk fills up. Large shuffles need local disk space on every node, not just in object storage.
spark.shuffle.sort.bypassMergeThreshold(200) lets RDD shuffles with no map-side aggregation and few reduce partitions write one file per reducer and concatenate them, skipping the sort. It is an internal optimisation, rarely something to tune.
In interviews
“What actually happens in a shuffle?” Describe map tasks partitioning and sorting their output into a data file and index file on local disk, reducers fetching their slice from every map output over the network, and the stage boundary between them. Mention that shuffle files outlive the stage (so retries and later jobs can reuse them, which is why the UI shows skipped stages).
Understanding shuffle cost
Where the time goes
A shuffle pays for the same data several times:
- Serialization and compression of every row on the map side.
- Local disk writes, plus extra spill writes if the sort buffer overflows.
- Network transfer: every reducer reads from every mapper, so
Mmap tasks andRreduce partitions meanM × Rblocks. 1,000 × 1,000 is a million small fetches. - Disk reads, decompression and deserialization on the reduce side.
- A barrier: no reduce task can start until all map tasks finish, so the slowest map task delays everything.
The volume that crosses the shuffle is what you control. A plan shows it: in a groupBy, Spark aggregates partially before the Exchange (partial_sum), so only one row per key per map task crosses the network.
Ways to make shuffles cheaper
| Technique | Why it helps |
|---|---|
| Filter and select columns before the wide operation | Fewer rows and narrower rows to serialize and move |
| Aggregate before joining when you can | A join on pre-aggregated data shuffles one row per key |
| Broadcast the small side of a join | The large side does not shuffle at all |
Reuse the partitioning: one repartition by a key used by several later operations |
Later joins and aggregations on the same key skip their Exchange |
| Bucketed tables for repeated joins on the same key | The shuffle is paid once at write time |
Avoid needless orderBy, distinct and repartition |
Each adds a full shuffle |
Reading shuffle volume
In the Stages tab, Shuffle Write on the map stage and Shuffle Read on the next stage show the bytes moved. In the SQL tab, the Exchange node shows data size and records. Compare them with the stage’s input size: a shuffle much larger than the input points to a join that multiplied rows or a missing filter.
In interviews
“Why are shuffles expensive and how do you reduce them?” List disk, network, serialization and the stage barrier, then give concrete reductions: filter early, prune columns, broadcast small tables, pre-aggregate, and reuse partitioning. Quantify with the UI’s shuffle read and write sizes rather than guessing.
Repartition versus coalesce
What they do
Both change the number of partitions; they differ in how.
repartition(n)performs a full shuffle with round-robin distribution, producingnpartitions of nearly equal size.repartition(n, "col")hash-partitions by the column instead. It can increase or decrease the count.coalesce(n)merges existing partitions without a shuffle, so it can only decrease the count. It is cheap but keeps whatever imbalance existed, and it changes the parallelism of the whole stage it belongs to.
print("coalesce(3): ", sizes(events.coalesce(3)))
print("repartition(3): ", sizes(events.repartition(3)))
print("repartition(4, user_id):", sizes(events.repartition(4, "user_id")))
coalesce(3): [25000, 37500, 37500]
repartition(3): [33335, 33332, 33333]
repartition(4, user_id): [9760, 70060, 9520, 10660]
coalesce(3) glued eight equal partitions into three unequal ones. repartition(3) balanced them through a shuffle. Hash partitioning by user_id put all 60,000 rows of the hot user in one partition: hashing never splits a key.
The plans show the difference: repartition adds an Exchange, coalesce does not.
events.repartition(4).explain()
events.coalesce(3).explain()
== Physical Plan ==
AdaptiveSparkPlan isFinalPlan=false
+- Exchange RoundRobinPartitioning(4), REPARTITION_BY_NUM, [plan_id=150]
+- Project [id#0L, CASE WHEN ((id#0L % 10) < 6) THEN 0 ELSE (id#0L % 5000) END AS user_id#1L, cast((id#0L % 50) as int) AS amount#2]
+- Range (0, 100000, step=1, splits=8)
== Physical Plan ==
Coalesce 3
+- *(1) Project [id#0L, CASE WHEN ((id#0L % 10) < 6) THEN 0 ELSE (id#0L % 5000) END AS user_id#1L, cast((id#0L % 50) as int) AS amount#2]
+- *(1) Range (0, 100000, step=1, splits=8)
The coalesce trap
Because coalesce adds no stage boundary, it reduces the parallelism of everything upstream in the same stage. df.filter(...).withColumn(...).coalesce(1).write... runs the filter and column logic in a single task. If the upstream work is heavy, repartition(1) is often faster: the heavy work runs in parallel and only the final write is single-threaded.
When to use which
| Situation | Choice |
|---|---|
| Reduce partitions after a big filter, before writing | coalesce(n) if the remaining partitions are reasonably even; repartition(n) if they are not |
| Increase parallelism before an expensive narrow step | repartition(n) |
| Cluster rows by key before a write partitioned by that key | repartition("date"), so each task writes few output folders |
| Range-ordered output, or sorted files by a column | repartitionByRange(n, "col") |
| Write exactly one file | coalesce(1) for small data; avoid for large data |
In interviews
“Repartition versus coalesce?” State that repartition shuffles and balances while coalesce merges without a shuffle and can only reduce; then add the trap that coalesce shrinks upstream parallelism within its stage. That nuance separates a memorised answer from an experienced one.
Choosing a partitioning strategy
The three partitioners
| Partitioning | How rows are placed | Produced by | Good for |
|---|---|---|---|
| Round robin | Rotates rows across partitions | repartition(n) |
Evening out sizes; no key needed |
| Hash | hash(key) mod n |
repartition(n, cols), groupBy, shuffle joins |
Co-locating equal keys for aggregation and joins |
| Range | Sorted key ranges from sampling | repartitionByRange, orderBy |
Sorted output, range filters, ordered files |
Spark tracks the partitioning of a plan. If a DataFrame is already hash-partitioned by user_id into the right number of partitions, a following groupBy("user_id") or join on user_id needs no new Exchange. You can see this by checking that the plan has only one Exchange.
Two meanings of “partitioning”
Do not mix up in-memory partitions with storage partitioning. df.write.partitionBy("event_date") creates one folder per date on disk, so later reads that filter on the date skip whole folders (partition pruning). It does not control how many in-memory partitions the job uses, and it multiplies files: each task writes one file per distinct value it holds. That is why you often repartition("event_date") before a partitioned write.
Choosing a key
- Use the key you join or aggregate on most, so the shuffle is reused.
- Check its cardinality and distribution. A low-cardinality key (country, status) gives few useful partitions; a skewed key (customer with a giant account) gives one huge one.
- For storage partitions, pick a column with modest cardinality that queries filter on, such as a date. A high-cardinality storage key (user id) creates millions of tiny files.
In interviews
“How would you partition this table?” Ask about the query patterns first. Then separate storage partitioning (folders for pruning, low-to-medium cardinality, usually date) from shuffle partitioning (hash by the join or group key, sized for tens to hundreds of megabytes per partition), and mention bucketing for repeated joins.
Tuning spark.sql.shuffle.partitions
What it is
spark.sql.shuffle.partitions sets how many partitions a DataFrame or SQL shuffle produces. It defaults to 200, a value that is too many for small data and too few for terabytes.
With Adaptive Query Execution (the default since Spark 3.2)
AQE coalesces small shuffle partitions after the map stage finishes, aiming at spark.sql.adaptive.advisoryPartitionSizeInBytes (64 MB by default) and never below spark.sql.adaptive.coalescePartitions.minPartitionSize (1 MB). So the setting becomes a starting point and upper bound. You can also set spark.sql.adaptive.coalescePartitions.initialPartitionNum to start higher for large jobs.
small = events.groupBy("amount").count()
spark.conf.set("spark.sql.adaptive.enabled", "false")
print("AQE off:", small.rdd.getNumPartitions())
spark.conf.set("spark.sql.adaptive.enabled", "true")
small = events.groupBy("amount").count()
print("AQE on: ", small.rdd.getNumPartitions())
AQE off: 8
AQE on: 1
With AQE off, 50 groups were spread over the configured 8 partitions. With AQE on, Spark saw that the whole shuffle was a few kilobytes and merged it into one partition.
One subtlety: spark.sql.adaptive.coalescePartitions.parallelismFirst defaults to true, which tells AQE to ignore the advisory size and only respect the 1 MB minimum, to keep as many tasks as possible running in parallel. The Spark docs recommend setting it to false on busy clusters so the advisory size is honoured.
Sizing it by hand
When AQE is off, or to set a sensible initial number for a large job:
- Find the shuffle write size of the biggest stage in the UI (say 500 GB).
- Divide by a target partition size (say 200 MB): 500,000 MB / 200 MB = 2,500 partitions.
- Round to a multiple of the total cores so the last wave is full (with 400 cores, 2,400 or 2,800).
Pitfalls
- Setting 2,000 globally for a pipeline that also runs small queries: without AQE each small query launches 2,000 tiny tasks.
- Forgetting that
spark.sql.shuffle.partitionsaffects DataFrame and SQL shuffles only. RDD operations usespark.default.parallelismor explicit arguments. - Changing the value for a streaming query after it has started: stateful streaming queries keep the partition count stored in their checkpoint.
In interviews
“How do you choose the number of shuffle partitions?” Answer with the arithmetic (shuffle size divided by a target of roughly 100–200 MB, rounded to the core count), then note that AQE coalescing makes the default far less harmful today, and say you would confirm with the task durations in the UI.
Detecting data skew
What it is
Skew means some keys have far more rows than others. After a shuffle by that key, one partition receives the hot key and its task runs much longer than the rest, so the whole stage waits for it.
Detect it in the data
A count per key, sorted descending, is the quickest check:
events.groupBy("user_id").count().orderBy(F.desc("count")).show(3)
+-------+-----+
|user_id|count|
+-------+-----+
| 0|60000|
| 26| 20|
| 28| 20|
+-------+-----+
only showing top 3 rows
You can also see where rows land after hash partitioning with spark_partition_id():
(events.repartition(8, "user_id")
.groupBy(F.spark_partition_id().alias("partition")).count()
.orderBy("partition").show())
+---------+-----+
|partition|count|
+---------+-----+
| 0| 4680|
| 1| 5180|
| 2| 4900|
| 3| 5280|
| 4| 5080|
| 5|64880|
| 6| 4620|
| 7| 5380|
+---------+-----+
Detect it in task metrics
In the UI’s stage detail, the summary metrics show min, median and max per task. Skew looks like a max far above the median for duration and shuffle read. The same numbers are available from the REST API. This runs a join with broadcasting and AQE turned off, then prints the records each reduce task read:
spark.conf.set("spark.sql.autoBroadcastJoinThreshold", "-1")
spark.conf.set("spark.sql.adaptive.enabled", "false")
def reduce_task_records(description):
for s in api("stages?details=true"):
if s.get("description") == description and s["shuffleReadBytes"] > 0 and s["numTasks"] == 8:
recs = [t["taskMetrics"]["shuffleReadMetrics"]["recordsRead"] for t in s["tasks"].values()]
if statistics.median(recs) > 0:
return sorted(recs)
sc.setJobDescription("skewed join")
print(events.join(users, "user_id").groupBy("country").count().orderBy("country").collect())
sc.setJobDescription(None)
recs = reduce_task_records("skewed join")
print("records per join task:", recs)
print("median:", statistics.median(recs), "max:", max(recs))
[Row(country='in', count=13340), Row(country='uk', count=73340), Row(country='us', count=13320)]
records per join task: [5216, 5270, 5506, 5709, 5820, 5950, 6054, 65475]
median: 5764.5 max: 65475
One join task read more than ten times the median. On a real cluster that task’s duration would dominate the stage.
Common causes
NULL, empty or default keys (-1,"unknown",0) produced by upstream bugs or optional fields.- Genuinely popular entities: a marketplace’s biggest seller, a viral post, a bot.
- Low-cardinality keys used for joins or partitioning.
- Time skew: one day or hour with far more data, such as a sale event.
In interviews
“How do you know a job is skewed?” Describe the symptom (a few tasks much slower; max far above median in stage metrics, often with spill on those tasks), then the confirmation (a count by key, or spark_partition_id() counts), and only then the fix.
Salting for skew
What it is
Salting splits a hot key into several artificial keys by adding a random number, the salt, so its rows spread across several partitions. It works on any Spark version and for both aggregations and joins, at the cost of extra code and some extra work.
Salting an aggregation (two-phase aggregation)
Aggregate by (key, salt) first, which spreads the hot key over up to SALT tasks, then aggregate the partial results by key alone:
SALT = 8
salted_totals = (events
.withColumn("salt", (F.rand(seed=42) * SALT).cast("int"))
.groupBy("user_id", "salt").agg(F.sum("amount").alias("partial"))
.groupBy("user_id").agg(F.sum("partial").alias("total")))
plain_totals = events.groupBy("user_id").agg(F.sum("amount").alias("total"))
print(salted_totals.filter("user_id = 0").collect())
print(plain_totals.filter("user_id = 0").collect())
print("rows that differ:", salted_totals.exceptAll(plain_totals).count())
[Row(user_id=0, total=1350000)]
[Row(user_id=0, total=1350000)]
rows that differ: 0
This works only for aggregations that can be combined in two steps: sum, count, min, max, and averages done as sum and count. A distinct count cannot simply be summed across salts.
DataFrame aggregations already combine partially on the map side, so a plain sum is rarely skewed badly. Salting aggregations matters most for functions without partial aggregation, such as collect_list, or for Python UDF aggregations.
Salting a join
For a join, salt the large, skewed side with a random salt and replicate the small side once per salt value, so every salted key still finds its match. Salting only the hot keys keeps the replication small:
salted_events = events.withColumn(
"salt", F.when(F.col("user_id") == 0, (F.rand(seed=7) * SALT).cast("int")).otherwise(F.lit(0)))
salted_users = users.withColumn(
"salt", F.explode(F.when(F.col("user_id") == 0, F.sequence(F.lit(0), F.lit(SALT - 1)))
.otherwise(F.array(F.lit(0)))))
print("users rows after replicating the hot key:", salted_users.count())
sc.setJobDescription("salted join")
print(salted_events.join(salted_users, ["user_id", "salt"])
.groupBy("country").count().orderBy("country").collect())
sc.setJobDescription(None)
recs = reduce_task_records("salted join")
print("records per join task:", recs)
print("median:", statistics.median(recs), "max:", max(recs))
users rows after replicating the hot key: 5007
[Row(country='in', count=13340), Row(country='uk', count=73340), Row(country='us', count=13320)]
records per join task: [5328, 5442, 5565, 13286, 13398, 13494, 20210, 28284]
median: 13342.0 max: 28284
Same result, and the largest task now reads far fewer records. With only 8 partitions some salts land together; with hundreds of partitions on a cluster the spread is more even.
Pitfalls
- Salting the wrong side. Replicate the side that is small for the hot key, and salt the large side.
- Too much salt. Replication multiplies the small side by the salt count; size it to the skew (enough to bring the hot key near the median partition size).
- Non-deterministic salts in retries.
rand()with a seed is deterministic per partition, which keeps results stable when a task is retried. Without a seed, salts can differ between attempts; that is still correct for a join or sum, but it can surprise you when comparing runs. - Salting when you could filter. If the hot key is
NULLor a default that never matches, filter it out or handle it separately instead.
In interviews
Be ready to write the two-phase aggregation and to describe the join version (random salt on the big side, explode the small side over all salts). Strong candidates mention the cost (replication, extra aggregation) and that AQE’s skew join often makes manual salting unnecessary for sort-merge joins.
Skew join optimisation
Options, in the order to try them
- Let AQE split skewed partitions. With
spark.sql.adaptive.skewJoin.enabled(on by default), AQE inspects the shuffle statistics of a sort-merge join. A partition is skewed when it is larger thanspark.sql.adaptive.skewJoin.skewedPartitionFactor(5) times the median partition size and larger thanspark.sql.adaptive.skewJoin.skewedPartitionThresholdInBytes(256 MB). AQE splits it into several tasks of about the advisory size and duplicates the matching partition of the other side for each. - Broadcast the small side. If one side fits under
spark.sql.autoBroadcastJoinThreshold(10 MB by default) or abroadcast()hint, the large side is never shuffled, so its skew does not matter. - Handle hot keys separately. Filter out
NULLand default keys that cannot match, or split the hot keys into their own join (often a broadcast join) andunionthe results. - Salt the key, as above.
- Fix the data: pre-aggregate, deduplicate, or choose a different join key upstream.
Seeing AQE’s skew join
On a laptop the data is far below 256 MB, so this example lowers the thresholds to a few kilobytes to trigger the optimisation. On a cluster you would normally keep the defaults.
spark.conf.set("spark.sql.adaptive.enabled", "true")
spark.conf.set("spark.sql.adaptive.coalescePartitions.enabled", "false")
spark.conf.set("spark.sql.adaptive.skewJoin.skewedPartitionThresholdInBytes", "4KB")
spark.conf.set("spark.sql.adaptive.advisoryPartitionSizeInBytes", "4KB")
by_country = events.join(users, "user_id").groupBy("country").agg(F.sum("amount").alias("amount"))
by_country.collect()
by_country.explain()
== Physical Plan ==
AdaptiveSparkPlan isFinalPlan=true
+- == Final Plan ==
ResultQueryStage 3
+- *(6) HashAggregate(keys=[country#5], functions=[sum(amount#2)])
+- ShuffleQueryStage 2
+- Exchange hashpartitioning(country#5, 8), ENSURE_REQUIREMENTS, [plan_id=1113]
+- *(5) HashAggregate(keys=[country#5], functions=[partial_sum(amount#2)])
+- *(5) Project [amount#2, country#5]
+- *(5) SortMergeJoin(skew=true) [user_id#1L], [user_id#4L], Inner
:- *(3) Sort [user_id#1L ASC NULLS FIRST], false, 0
: +- AQEShuffleRead skewed
: +- ShuffleQueryStage 0
: +- Exchange hashpartitioning(user_id#1L, 8), ENSURE_REQUIREMENTS, [plan_id=993]
: +- *(1) Project [CASE WHEN ((id#0L % 10) < 6) THEN 0 ELSE (id#0L % 5000) END AS user_id#1L, cast((id#0L % 50) as int) AS amount#2]
: +- *(1) Filter CASE WHEN ((id#0L % 10) < 6) THEN true ELSE isnotnull((id#0L % 5000)) END
: +- *(1) Range (0, 100000, step=1, splits=8)
+- *(4) Sort [user_id#4L ASC NULLS FIRST], false, 0
+- AQEShuffleRead
+- ShuffleQueryStage 1
+- Exchange hashpartitioning(user_id#4L, 8), ENSURE_REQUIREMENTS, [plan_id=1000]
+- *(2) Project [id#3L AS user_id#4L, element_at([uk,in,us], cast(((id#3L % 3) + 1) as int), None, true) AS country#5]
+- *(2) Range (0, 5000, step=1, splits=2)
+- == Initial Plan ==
(trimmed: the same plan without query stages or skew handling)
SortMergeJoin(skew=true) and AQEShuffleRead skewed show that AQE split the hot partition on the events side and replicated the matching users partition. In the UI the join stage shows more tasks than shuffle partitions.
Broadcast avoids the problem
When the dimension is small, a broadcast join is simpler still: there is no Exchange on the large side, so nothing to be skewed.
for k in ["spark.sql.adaptive.coalescePartitions.enabled",
"spark.sql.adaptive.skewJoin.skewedPartitionThresholdInBytes",
"spark.sql.adaptive.advisoryPartitionSizeInBytes",
"spark.sql.autoBroadcastJoinThreshold"]:
spark.conf.unset(k)
events.join(F.broadcast(users), "user_id").explain()
== Physical Plan ==
AdaptiveSparkPlan isFinalPlan=false
+- Project [user_id#1L, id#0L, amount#2, country#5]
+- BroadcastHashJoin [user_id#1L], [user_id#4L], Inner, BuildRight, false, false
:- Project [id#0L, CASE WHEN ((id#0L % 10) < 6) THEN 0 ELSE (id#0L % 5000) END AS user_id#1L, cast((id#0L % 50) as int) AS amount#2]
: +- Filter CASE WHEN ((id#0L % 10) < 6) THEN true ELSE isnotnull((id#0L % 5000)) END
: +- Range (0, 100000, step=1, splits=8)
+- BroadcastExchange HashedRelationBroadcastMode(List(input[0, bigint, false]),false), [plan_id=1168]
+- Project [id#3L AS user_id#4L, element_at([uk,in,us], cast(((id#3L % 3) + 1) as int), None, true) AS country#5]
+- Range (0, 5000, step=1, splits=2)
What AQE skew handling does not cover
- Skewed aggregations. AQE splits join partitions, not
groupBypartitions. Map-side partial aggregation usually saves you; for non-combinable aggregations, salt. - Outer joins on the wrong side. AQE can split the left side of a left outer join but not the right side (the side whose unmatched rows must be preserved cannot be split), and similar rules apply to other join types.
- Joins where splitting would require an extra shuffle. AQE skips the optimisation unless
spark.sql.adaptive.forceOptimizeSkewedJoinis set. - Skew inside one task that is not about partition size, such as a single enormous row or an expensive UDF on some values.
In interviews
“How do you fix a skewed join?” Give an ordered answer: confirm it in the UI, check whether AQE’s skew join applied (look for skew=true in the final plan in the SQL tab), broadcast if one side is small, isolate or filter hot keys, and salt as the general fallback. Mention the AQE thresholds (factor 5 and 256 MB) and the limits above.
Practice questions
A 2 TB dataset is read and aggregated. The job has 200 reduce tasks and some spill to disk. What do you change?
The shuffle partitions are too large: roughly 2 TB / 200 is 10 GB each (less after partial aggregation, but still big). Raise spark.sql.shuffle.partitions, or spark.sql.adaptive.coalescePartitions.initialPartitionNum with AQE, to aim for roughly 100–200 MB per partition, for example 10,000 to 20,000, and let AQE coalesce down where the data is smaller. Check the spill and task durations again afterwards.
Why can coalesce(1) before a write make a job much slower than repartition(1)?
coalesce does not add a shuffle, so the single partition applies to the whole stage containing it: every upstream narrow transformation runs in one task. repartition(1) adds a stage boundary, so upstream work runs with full parallelism and only the final write runs in one task.
Describe the files a shuffle writes and how reducers read them.
Each map task writes one data file sorted by reduce partition id, plus an index file of byte offsets (and a checksum file), to local disk under spark.local.dir. Each reducer asks the driver where map outputs are, then fetches its byte range from every map task’s data file over the network, a limited amount in flight at a time, and merges them. The files outlive the stage, so retries and later jobs can reuse them.
Stage metrics show median task duration 10 s and max 25 min, with 40 GB shuffle read on the slow task. What are your next steps?
That is skew. Find the hot key with a count per key on the join or group key. If it is a NULL or default value that cannot match, filter it or handle it separately. If it is a join, check the final plan for AQE’s skew=true; if AQE did not split it (wrong join side, thresholds, extra shuffle needed), consider broadcasting the smaller side, separating the hot key into its own join, or salting. Re-run and compare max versus median again.
When does AQE treat a join partition as skewed?
When it is larger than skewedPartitionFactor (default 5) times the median partition size and also larger than skewedPartitionThresholdInBytes (default 256 MB), with spark.sql.adaptive.skewJoin.enabled on (the default). It then splits the partition into roughly advisory-size pieces and duplicates the other side’s matching partition.
Write the idea of a two-phase (salted) aggregation and say when it is needed.
Add a random salt column in a range 0 to N−1, aggregate by (key, salt) to get partial results, then aggregate the partials by key. It spreads a hot key across N tasks. It is needed when the aggregation cannot be combined on the map side by Spark (for example collect_list or a Python aggregation UDF) or when a single key’s partial results are still too large; ordinary sum and count already get partial aggregation.
What is the difference between df.repartition(“date”) and df.write.partitionBy(“date”)?
repartition("date") is an in-memory hash shuffle, so all rows of a date go to the same task. write.partitionBy("date") lays out output folders by date for partition pruning on read. Using both together means each task writes files for only a few dates, avoiding many small files.
Key takeaways
- One task per partition per stage; total cores cap concurrency, so aim for several partitions per core and even partition sizes.
- A shuffle writes partitioned, sorted files on every mapper and has every reducer fetch from every mapper; it costs disk, network, serialization and a stage barrier.
repartitionshuffles and balances;coalescemerges without a shuffle and also shrinks upstream parallelism in its stage.spark.sql.shuffle.partitions(200) is a starting point under AQE, which coalesces towards 64 MB partitions; size it from shuffle volume for large jobs.- Detect skew from max versus median task metrics and from a count per key, before choosing a fix.
- Fix skewed joins with AQE skew join, broadcasting, isolating hot keys or salting; AQE does not fix skewed aggregations.
Progress is saved in this browser only. No account needed.