Menu

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.

  • Intermediate
  • 28 min read
  • Updated Oct 2026
On this page
  1. Sample data
  2. Partitions and parallelism
  3. What it is
  4. Where partition counts come from
  5. How many partitions you want
  6. Pitfalls
  7. In interviews
  8. Shuffle internals
  9. What it is
  10. The map side: write sorted, partitioned files
  11. The reduce side: fetch and merge
  12. Who serves the files
  13. Pitfalls
  14. In interviews
  15. Understanding shuffle cost
  16. Where the time goes
  17. Ways to make shuffles cheaper
  18. Reading shuffle volume
  19. In interviews
  20. Repartition versus coalesce
  21. What they do
  22. The coalesce trap
  23. When to use which
  24. In interviews
  25. Choosing a partitioning strategy
  26. The three partitioners
  27. Two meanings of “partitioning”
  28. Choosing a key
  29. In interviews
  30. Tuning spark.sql.shuffle.partitions
  31. What it is
  32. With Adaptive Query Execution (the default since Spark 3.2)
  33. Sizing it by hand
  34. Pitfalls
  35. In interviews
  36. Detecting data skew
  37. What it is
  38. Detect it in the data
  39. Detect it in task metrics
  40. Common causes
  41. In interviews
  42. Salting for skew
  43. What it is
  44. Salting an aggregation (two-phase aggregation)
  45. Salting a join
  46. Pitfalls
  47. In interviews
  48. Skew join optimisation
  49. Options, in the order to try them
  50. Seeing AQE’s skew join
  51. Broadcast avoids the problem
  52. What AQE skew handling does not cover
  53. In interviews
  54. Practice questions
  55. 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.range or parallelize without a partition count in tests: you get defaultParallelism, 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:

  1. Serialization and compression of every row on the map side.
  2. Local disk writes, plus extra spill writes if the sort buffer overflows.
  3. Network transfer: every reducer reads from every mapper, so M map tasks and R reduce partitions mean M × R blocks. 1,000 × 1,000 is a million small fetches.
  4. Disk reads, decompression and deserialization on the reduce side.
  5. 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, producing n partitions 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:

  1. Find the shuffle write size of the biggest stage in the UI (say 500 GB).
  2. Divide by a target partition size (say 200 MB): 500,000 MB / 200 MB = 2,500 partitions.
  3. 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.partitions affects DataFrame and SQL shuffles only. RDD operations use spark.default.parallelism or 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 NULL or 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

  1. 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 than spark.sql.adaptive.skewJoin.skewedPartitionFactor (5) times the median partition size and larger than spark.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.
  2. Broadcast the small side. If one side fits under spark.sql.autoBroadcastJoinThreshold (10 MB by default) or a broadcast() hint, the large side is never shuffled, so its skew does not matter.
  3. Handle hot keys separately. Filter out NULL and default keys that cannot match, or split the hot keys into their own join (often a broadcast join) and union the results.
  4. Salt the key, as above.
  5. 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 groupBy partitions. 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.forceOptimizeSkewedJoin is 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.
  • repartition shuffles and balances; coalesce merges 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.

By Data Career Hub Editorial · Last reviewed Oct 2026 · All examples run on PySpark 4.2.0 in local mode (local[2]). Shuffle partitions are set to 8 and some AQE thresholds are lowered to a few kilobytes so effects that need gigabytes on a cluster appear on a laptop; the lesson says where. Configuration defaults were read from the Spark 4.2 runtime.

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

Search
Filter by type