A join in Spark is usually a shuffle: both sides are repartitioned by the hash of the join key so that matching rows meet in the same task. That works as long as keys are roughly the same size. When one key owns a billion rows, every one of them hashes to the same partition, and one task does the work that two hundred others share. The stage takes as long as that one task, it may spill or run out of memory, and adding executors changes nothing.

Salting fixes this by changing the key. You append a small number, the salt, to the join key on the big side so the hot key's rows spread over several partitions, and you copy the matching rows on the other side once per salt value so every piece still finds its match. The idea fits in one sentence, but the details decide whether it works: how many salts, which keys, which side gets copied, and which join types survive the copying. This page works through them on a concrete example. For finding skew in the first place and the wider set of remedies, start with Spark data skew handling.

Advertisement

What salting changes, and what it costs

Without salting, the shuffle sends a row to partition hash(key) mod P. With salting, the big side uses hash(key, salt) mod P, where the salt is a number from 0 to n - 1 assigned per row, so a key with n salts is processed by up to n tasks.

The other side must be able to match any of those salts, so each of its rows for that key is copied n times, once with each salt value. That copying is the whole cost of salting. If the copied side is a dimension table with one row per key, copying the hot keys' rows is almost free. If you salt every key with the same n, you copy the entire table n times and turn a skew problem into a shuffle-volume problem. The rest of this page is about keeping that copy small.

Selective salting: split the hot key, replicate only its dimension rowpayments (fact)merchant_17: 1.9B rowsmerchants (dim)one row per merchanthot-key tablekey, n_salts (broadcast)joinjoinfact + saltsalt = hash(payment_id) mod ndim explodedn copies, salt 0..n-1shuffle on (key, salt)76 tasks for merchant_17Cold keys keep salt 0 and one dimension copyso the extra cost is only (n - 1) rows per hot key, not N times the whole table
Selective salting. A small hot-key table, broadcast to every executor, decides how many salts each key gets. The fact side hashes a per-row id into one salt; the dimension side is copied once per salt for hot keys only.

When AQE is enough, and when it is not

Since Spark 3.0, adaptive query execution can handle skewed sort-merge joins on its own. After the shuffle map stage finishes, Spark knows the real size of every shuffle partition. A partition counts as skewed when it is larger than spark.sql.adaptive.skewJoin.skewedPartitionFactor (default 5.0) times the median partition and also larger than spark.sql.adaptive.skewJoin.skewedPartitionThresholdInBytes (default 256MB). Spark then splits that partition into pieces of about spark.sql.adaptive.advisoryPartitionSizeInBytes (default 64 MB) and reads the matching partition on the other side once per piece, which is salting done for you at runtime. The rule is on by default through spark.sql.adaptive.skewJoin.enabled; the AQE architecture page explains the machinery.

Try AQE first. Reach for manual salting when one of these holds:

  • The split is too coarse. AQE splits a skewed partition along the boundaries of map output blocks: each piece is a range of mapper outputs. If the hot key's rows come from a few large map tasks, the pieces cannot get smaller than one mapper's contribution, and the straggler remains.
  • The rule declines to run. If splitting would require an extra shuffle, for example because a later operator depends on the join's output partitioning, AQE skips the optimization unless spark.sql.adaptive.forceOptimizeSkewedJoin is true. It also does not apply to every join type; check the plan rather than assuming.
  • The skew is not in a join's shuffle. A skewed window function, a skewed groupBy with a non-combinable aggregate, or a skewed write is not covered by the join rule.
  • The other side is small. If the non-skewed side fits in memory, broadcast it and skip the shuffle altogether; see broadcast joins. Salting is for when both sides are too big to broadcast.

To tell which case you are in, compare the stage's maximum task duration and shuffle read with the median in the Spark UI, and check in the SQL tab whether the join's shuffle read treated any partitions as skewed.

Advertisement

Step 1: measure the hot keys

Salting starts with numbers, not guesses. Count rows per join key on the big side and look at both the head and the body of the distribution. The aggregation is cheap because Spark combines counts on the map side before shuffling them.

from pyspark.sql import functions as F

facts = spark.table("payments")     # ~4.1B rows; join key merchant_id, unique payment_id
dims = spark.table("merchants")     # ~2M rows, one per merchant_id

key_sizes = facts.groupBy("merchant_id").count()

# The head of the distribution: the keys that will become stragglers.
key_sizes.orderBy(F.desc("count")).show(20, truncate=False)

# The body: what a normal key looks like.
key_sizes.select(F.percentile_approx("count", [0.5, 0.99], 1000)).show(truncate=False)

In the running example, a payments table of about 4.1 billion rows joins a merchants table of about 2 million rows. Three merchants have 1.9 billion, 610 million and 240 million rows; the 99th percentile key has about 40,000. A few enormous keys and a long normal tail is the common shape, and why salting every key is wasteful.

Also count null keys: they match nothing in an inner join, yet all hash to one partition. Filter them out before the join rather than salting them.

Step 2: choose the salt count per key

Choose n so that each salted piece is about the size you want one task to process. Pick that target from your own stages: look at a healthy task's input rows and duration, and choose a size that finishes in a minute or two without spilling. Then n for each hot key is its row count divided by the target, rounded up. Keys below the target get no salt.

import math

# Measured with a groupBy on the join key (see the query above).
HOT_KEYS = {"merchant_17": 1_900_000_000, "merchant_4": 610_000_000, "merchant_92": 240_000_000}
TYPICAL_KEY_ROWS = 40_000          # p99 of the non-hot keys
TARGET_ROWS_PER_TASK = 25_000_000  # what one task handles in about a minute here
DIM_ROWS_PER_KEY = 1               # one dimension row per merchant
ALL_KEY_SALT = 32                  # what blanket salting would use for every key

total_extra = 0
print(f"{'key':12} {'rows':>14} {'salts':>6} {'rows/salt':>12} {'extra dim rows':>15}")
for key, rows in HOT_KEYS.items():
    n = max(1, math.ceil(rows / TARGET_ROWS_PER_TASK))
    total_extra += (n - 1) * DIM_ROWS_PER_KEY
    print(f"{key:12} {rows:14,} {n:6d} {rows // n:12,} {(n - 1) * DIM_ROWS_PER_KEY:15,}")

dim_keys = 2_000_000
print(f"selective salting adds {total_extra:,} dimension rows")
print(f"blanket salting at N={ALL_KEY_SALT} adds {(ALL_KEY_SALT - 1) * dim_keys:,} dimension rows")

Running it prints:

key                    rows  salts    rows/salt  extra dim rows
merchant_17   1,900,000,000     76   25,000,000              75
merchant_4      610,000,000     25   24,400,000              24
merchant_92     240,000,000     10   24,000,000               9
selective salting adds 108 dimension rows
blanket salting at N=32 adds 62,000,000 dimension rows

Per-key counts keep every piece near the target, where a global n = 32 would leave merchant_17 in 59-million-row pieces. And because only three keys are salted, the dimension side grows by 108 rows instead of 62 million.

Make sure the shuffle has room for the pieces. With 111 salted pieces plus the cold keys, spark.sql.shuffle.partitions should be well above 111, or AQE's coalescing must be left on to merge the many small cold partitions afterwards. Spark partitioning covers that arithmetic.

Step 3: selective salting in PySpark

The implementation joins both sides to a small hot-key table, which is broadcast so that this extra join adds no shuffle. Cold keys get n_salts = 1, which turns their salt into the constant 0 and their copy count into one.

TARGET = 25_000_000

hot = (key_sizes
       .where(F.col("count") > TARGET)
       .select("merchant_id",
               F.ceil(F.col("count") / TARGET).cast("int").alias("n_salts")))
hot = hot.cache()
hot.count()                                   # a handful of rows; materialize once

def with_salt_count(df):
    return (df.join(F.broadcast(hot), "merchant_id", "left")
              .withColumn("n_salts", F.coalesce("n_salts", F.lit(1))))

# Fact side: every row gets exactly one salt. Hash a column that varies
# WITHIN the hot key (payment_id), never the join key itself.
facts_s = (with_salt_count(facts)
           .withColumn("salt", F.expr("pmod(xxhash64(payment_id), n_salts)").cast("int"))
           .drop("n_salts"))

# Dimension side: a hot key's row is copied once per salt; cold keys keep one copy.
dims_s = (with_salt_count(dims)
          .withColumn("salt", F.explode(F.sequence(F.lit(0), F.col("n_salts") - 1)))
          .drop("n_salts"))

joined = facts_s.join(dims_s, ["merchant_id", "salt"], "inner").drop("salt")
joined.explain()     # the shuffle keys should now read merchant_id, salt

Two details matter. The fact side's salt is a hash of payment_id, which differs between rows of the same merchant; hashing the join key would give every row of merchant_17 the same salt. And sequence(0, n - 1) with explode creates exactly n copies, so every fact row finds exactly one match. Key sizes drift, so persist the hot-key table and refresh it daily.

Which join types survive salting

Salting relies on one invariant: each row on the salted side appears exactly once, and each row on the copied side appears once per salt. Whether the join's output is still correct depends on what the join does with rows that do not match.

JoinSalt this sideCopy this sideNotes
Innerthe skewed sidethe other sidecorrect either way round
Left outerleft (preserved)righteach left row matches its one copy or none; output unchanged
Right outerright (preserved)leftmirror of left outer
Left semi / antileftrighta left row matches iff its key exists; copies cannot double it
Full outer——not safe: an unmatched row on the copied side appears n times in the output

The rule: the side whose unmatched rows are kept must be salted, never copied. A full outer join keeps both sides' unmatched rows, so split it instead: hot keys, which exist on the fact side by definition, need only a salted left outer join; run an ordinary full outer join on the cold keys and union the two.

When both sides are hot: two-dimensional salting

In a many-to-many join the hot key can be large on both sides, and copying one side n times copies hundreds of millions of rows. Split both instead: side A gets a salt i in 0..a-1 and is copied over every j in 0..b-1; side B gets a salt j and is copied over every i, and the join key becomes (key, i, j). Every pair of rows meets in exactly one of the a times b cells. Salting spreads the output over more tasks but does not shrink it, so aggregate first if the query does not need the product.

Deterministic salts

Many examples assign salts with rand(). Its output depends on the order in which a task sees its rows, and after a shuffle that order can differ when a task is retried, which costs recomputation during failures and makes results hard to reproduce. A hash of a stable per-row column gives the same salt on every run. Without a unique column, hash a nearly unique combination such as timestamp and user id.

Worked example: before and after

Before salting, most join tasks of the payments job read a few hundred megabytes and finish in under a minute; the task holding merchant_17 reads 1.9 billion rows, spills and runs for most of the job. AQE marks the partition as skewed, but the table was written by a handful of large upstream tasks, so the split pieces are still several hundred million rows each.

After selective salting, the stage has 111 salted pieces near 25 million rows each, plus the cold keys, and the dimension side grew by 108 rows. Check that the longest task is now within a small factor of the median. If not, the target is too large or a newly hot key is missing from the hot-key table. Compare runs on the same input, because skew moves from day to day.

Failure modes

  • Salting the join key's own hash: every row of the hot key gets the same salt and the straggler stays exactly where it was.
  • Copying the preserved side of an outer join: unmatched rows multiply silently. Row counts are the cheapest test; compare them with the unsalted join on a sample.
  • Blanket salting: a global n copies the whole dimension n times and can make the job slower than the skew did.
  • A stale hot-key table: new hot keys arrive unsalted and become the new stragglers.
  • Too few shuffle partitions: salted pieces collide in the same partitions and recreate the skew at a smaller scale.
  • Fighting the planner: a broadcast hint or AQE conversion may already avoid the shuffle; salting then only adds work. Read the plan first, and see join strategy hints for how Spark chooses.

Trade-offs

AQE's split is free but bounded by map output granularity; salting is code you must keep correct as data changes. Selective salting is cheap at runtime but needs a refreshed hot-key table. Prefer the simplest option that brings the longest task near the median.

What to do next

  1. Open the slow join stage and compare maximum with median task duration and shuffle read; confirm the straggler is a join key.
  2. Check whether AQE already marked the partition as skewed, and whether broadcasting the smaller side is possible.
  3. Count rows per join key, record the top keys and the 99th percentile, and filter null keys.
  4. Choose a target rows-per-task from a healthy stage and compute a salt count per hot key.
  5. Implement selective salting with a broadcast hot-key table and a salt hashed from a per-row id.
  6. Confirm the join type is safe to salt, and compare output row counts with the unsalted join on a sample.
  7. Schedule the hot-key table to refresh, and alert when the longest task exceeds a few times the median.
Key takeaway: Salting splits a hot join key into n pieces on the big side and copies the matching rows n times on the other side, so the copy is the cost to manage. Try AQE's skew split and broadcasting first; salt when map output granularity or the planner's rules leave a straggler. Measure key sizes, give each hot key its own salt count from a target rows per task, salt only those keys through a broadcast hot-key table, hash a per-row id rather than calling rand, and never copy the preserved side of an outer join.