Slow join in pyspark, tried repartition

Viewed 3816

I'm trying to left join 2 tables on Spark 3, with 17M rows (events) and 400M rows (details). have an EMR cluster of 1 + 15 x 64core instances. (r6g.16xlarge tried with similar r5a) Source files are unpartitioned parquet loaded from S3.

this is the code I'm using to join:

join = (
    broadcast(events).join(
        details,
        [
            details["a"] == events["a2"],
            (unix_timestamp(events["date"]) - unix_timestamp(details["date"])) / 3600
            > 5,
        ],
        "left",
    )
).drop("a")

join.checkpoint()

To partition I'm using this:

executors = 15 * 64 * 3  # 15 instances, 64 cores, 3 workers per core

so I tried:

details = details.repartition(executors, "a")

and

details = details.withColumn("salt", (rand(seed=42) * nSaltBins).cast("int"))
details = details.repartition(executors, "salt")

In both scenarios, 90% of the workers end in around 5-10 minutes and the rest continue for a LONG time (50+minutes), long green line, no memory or disk errors on the log.

There is a little skewness after partitioning (all partitions between 180k and 160k rows), nothing accountable for more than 50 minutes of processor time.

Any idea of what I could be overseeing? Read a ton of posts and still feel that the green lines (worker time) should be closer between each other, they are all starting at the same time, they are not waiting for a worker to end.

Thanks!

--Edit--- Removed broadcast

On job 11, stage 17 it does 974/1000 in 2 minutes and 30 min later still on 993/1000, previous step uses the salted partitions (given by the executors variable) and it's very fast.

Execution plan:

Using 17906254 events
== Physical Plan ==
AdaptiveSparkPlan (13)
+- Project (12)
   +- SortMergeJoin LeftOuter (11)
      :- Sort (4)
      :  +- Exchange (3)
      :     +- Project (2)
      :        +- Scan parquet  (1)
      +- Sort (10)
         +- Exchange (9)
            +- Exchange (8)
               +- Project (7)
                  +- Filter (6)
                     +- Scan parquet  (5)

as image

An example of 2h and more of 25% of that time is 1 executor remaining enter image description here

Current spark configuration:

spark = SparkSession.builder.appName('Test').config("spark.driver.memory", "108g").config(
        "spark.executor.instances", "59").config("spark.executor.memoryOverhead", "13312").config(
        "spark.executor.memory", "108g").config("spark.executor.cores", "15").config("spark.driver.cores", "15").config(
        "spark.default.parallelism", "1770").config("spark.sql.adaptive.enabled", "true").config(
        "spark.sql.adaptive.skewJoin.enabled", "true").config("spark.sql.shuffle.partitions", "885").getOrCreate()
2 Answers

Your issue looks like a nice case of skewed join where some partition will get a lot more data than the others and thus slow the complete job.

Repartitioning your dataframe before your join will not help because the SortMergeJoin operation will repartition again on your join keys to process the join

Since you're using Spark 3, you should have support for the automatic skewJoin management.

To use it, make sure you have both spark.sql.adaptive.enabled=true (it's false by default in standard Spark distribution) and spark.sql.adaptive.skewJoin.enabled=true

If you can't use automatic skewJoin optimization, you can fix it manually with something like this:

  • duplicate N times your small dataset
n = 10   # Chose an appropriate amount based on skewness
skewedEvents = events.crossJoin(spark.range(0,n).withColumnRenamed("id","eventSalt"))
  • seed your large dataset with a random column value between 0 and N
import pyspark.sql.functions as f

skewedDetails = details.withColumn("detailSalt", (f.rand() * n).cast("int"))
  • join using your salt in the join key and then drop the salt
joined = skewedEvents.join(skewedDetails,[         [
            skewedDetails["a"] == skewedEvents["a2"],
            skewedDetails["detailSalt"] == skewedEvents["eventSalt"],
            (unix_timestamp(skewedEvents["date"]) - unix_timestamp(skewedDetails["date"])) / 3600
            > 5,
        ],
        "left")\
        .filter("a is not null or (a is null and eventSalt = 0)")\
        .drop("a").drop("eventSalt").drop("detailSalt")

Note that also may want to validate your query join condition because the UI shows that with 333 million rows processed on details and 17 million on events you generated more than 5 billions output rows, so you may match more rows that you think with your join condition.

broadcast() is used to cache data on each executor (instead of sending the data with every task) but it's not working too well with very large amounts of data. It seems here that 17M rows was a bit too much.

Pre-partitionning your source data before the join could also help if the partitioning of the source data is not optimized for the join. You'll want partition around the column you use for the join. Usually data should be partitionned depending on how it's consumed.

Related