SpiralTrain
Exercises › Block 2 · Exercise 7

Block 2 · Exercise 7

Performance

Open the notebook 19-performance-starter in your workspace and run the setup cell at the top. It starts the Spark session, which takes a few minutes ; read Step 1 while it starts. Steps 4, 6 and 7 have a cell with None where your code goes, Step 8 is an open task on a slow query, and the other steps run as they stand.

Every timing you take is your own. With the whole class working at once, a query runs slower than on the quiet capacity where the slides were measured. Compare your own runs with each other, and with nobody else's numbers.

Do one part at a time, 1a, 1b and so on, and compare with the expected output before you go on. The solution is in 19-performance-solution, under every part on the exercise site, and at the back of the exercises PDF.

Step 1Tables

The setup cell loads four tables :

  • flights from flights_large : the real flights of the previous modules, for more months, stored in one folder per year and month
  • airports : the OurAirports table
  • pings and receivers : generated tables for Step 4, joined on receiver_id. pings has one row per position report from an aircraft, receivers one row per ground receiver

1a

Run the cell of Step 1. It counts the flights per month. How many months are there, and roughly how many flights in each?

Check your output
+----+-----+------+
|year|month| count|
+----+-----+------+
| ...| ... |   ...|
+----+-----+------+

One row per month, in order. The counts depend on the size of the table you have.

Step 2Busiest airport

2a

Run the cell of Step 2. It answers three questions with one query each. Write down how many flights, carriers and departure airports flights holds, and how many airports in airports have at least one departure.

Check your output
+-------+--------+------------------+
|flights|carriers|departure_airports|
+-------+--------+------------------+
|    ...|     ...|               ...|
+-------+--------+------------------+

airports with at least one departure : ...

2b

The last query of the cell shows the five airports with the most departures and their share of all flights. Which airport has the most departures, and what share of all flights is that? Write both down : it is the skew in this data, and Step 4 comes back to it.

Check your output
+------+------+-----+
|origin| count|share|
+------+------+-----+
|   ...|   ...|  ...|
+------+------+-----+

Five rows, the busiest first.

Step 3Distinct routes

A route is an origin and a dest together. The cells of Step 3 count the distinct routes flown on every day.

3a

Run the first cell of Step 3. It shows the ten days with the most routes, then prints the formatted plan.

Check your output
+-----------+------+
|flight_date|routes|
+-----------+------+
|        ...|   ...|
+-----------+------+

== Physical Plan ==
AdaptiveSparkPlan ...

Ten rows, with routes falling from the top. The plan lists its nodes, and its shape depends on the Spark version.

3b

Read the plan before you run anything else. How many Exchange nodes are there, and how many HashAggregate nodes? Why more than one HashAggregate? Count them in code and write the answer to the last question in a comment.

Hint 1

The slide Reading a Query Plan shows how the nodes are listed. The first block of the formatted output is the tree.

Check your output
Exchange nodes      : ...
HashAggregate nodes : ...
Show solutionHide solution
python
tree = plan_text(routes_per_day).split(chr(10) + chr(10))[0]
print("Exchange nodes      :", tree.count(" Exchange ("))
print("HashAggregate nodes :", tree.count(" HashAggregate ("))

# An aggregation runs twice : a partial one inside every partition before the shuffle,
# and a final one after it that combines the partial results. Counting distinct routes
# adds a step for the distinct pairs, so there are more than two.
# Every Exchange is a shuffle : the aggregation's and the sort's.

3c

Run the second cell of Step 3. It runs the query with collect() and prints the plan again. What does the first line say now? Look for AQEShuffleRead.

Hint 1

The slide Adaptive Query Execution says when AQE chooses its plan.

Check your output
== Physical Plan ==
AdaptiveSparkPlan isFinalPlan=true
...
   +- AQEShuffleRead ...

AQEShuffleRead appears where AQE merged small partitions after a shuffle.

3d

Would show() in place of collect() have given the final plan? Write why in a comment.

Hint 1

show() does not run the DataFrame's own plan again. What does it run?

Show solutionHide solution
python
# No. show() runs a new query of its own, so the DataFrame's plan stays the initial one,
# isFinalPlan=false. The SQL tab of the Spark UI shows the final plan either way.

Step 4Skewed join

The real flights are not skewed enough to show a hot key. From 4c on, the join uses pings, a generated table in which receiver 1 hears nine pings in ten. It is made up on purpose : it is the clearest way to see what a hot key does. Every callsign in it starts with SIM.

4a

Run the first cell of Step 4. It turns automatic broadcasting off, so Spark has to shuffle both sides, joins every flight to its departure airport and averages the departure delay per airport type.

Check your output
+--------------+---------+
|          type|avg_delay|
+--------------+---------+
|           ...|      ...|
+--------------+---------+

One row per airport type.

4b

Open the Spark UI and find the stage that performed the join. In its summary metrics, compare the maximum task duration with the median. How far apart are they? Which airport do you expect the slowest task held?

Hint 1

Look at the answer to 2b.

Show solutionHide solution
python
# Expected : the busiest airport of Step 2. Every row of one key hashes to the same partition,
# so the task that holds the hot key has the most rows and takes the longest.

4c

Run the second cell of Step 4. It switches AQE's skew handling off, so that only your own fix can help. It selects only the columns the join and its result need : a shuffle moves every column it is given. It joins pings to receivers on receiver_id and times the join with a noop write. Write down the time, and the maximum and the median task duration of the join stage in the Spark UI.

Check your output
plain join : ... s

The time is your own.

4d

Salt the join, for the hot receiver only. The key is an integer, so the salt goes in a second column and not on the end of the key. Fill in the cell with salted_pings and salted_receivers :

  • on ping_cols, add a column salt : a random number from 0 to 7 on the rows of receiver 1, and 0 on every other row
  • on receiver_cols, replicate the row of receiver 1 eight times, once per salt value, and give every other receiver salt 0
Hint 1

The slide Salting a Join shows the idea on a different table. F.when(...).otherwise(...) chooses per row.

Hint 2

F.rand(seed=7) gives a number between 0 and 1. Multiply and cast to int. F.explode of an array turns one row into as many rows as the array has items.

Show solutionHide solution
python
HOT = 1
SALTS = 8

salted_pings = ping_cols.withColumn(
    "salt", F.when(F.col("receiver_id") == HOT, (F.rand(seed=7) * SALTS).cast("int")).otherwise(0))
salted_receivers = receiver_cols.withColumn(
    "salt", F.explode(F.when(F.col("receiver_id") == HOT, F.array(*[F.lit(i) for i in range(SALTS)]))
                       .otherwise(F.array(F.lit(0)))))

4e

Join the two salted tables on both columns, ["receiver_id", "salt"], and drop the salt afterwards. Print the number of rows in receiver_cols and in salted_receivers. Time the salted join with a noop write, as in 4c, and check that it gives as many rows as the plain join.

Hint 1

.drop("salt") after the join keeps the result the same shape as the plain join.

Check your output
receivers : ... rows ; salted : ... rows
salted join : ... s
same rows : True

The salted receivers hold seven more rows than the receivers. The time is your own.

Show solutionHide solution
python
salted_join = salted_pings.join(salted_receivers, ["receiver_id", "salt"]).drop("salt")

print("receivers :", receiver_cols.count(), "rows ; salted :", salted_receivers.count(), "rows")
start = time.time()
salted_join.write.format("noop").mode("overwrite").save()
print(f"salted join : {time.time() - start:.1f} s")
print("same rows :", salted_join.count() == ping_cols.join(receiver_cols, "receiver_id").count())

4f

Open the join stage of the salted join in the Spark UI. Did the maximum task duration fall? Did the total time fall, or did the extra rows eat the gain? Compare your times from 4c and 4e and write the answer in a comment.

Show solutionHide solution
python
# There is no number to copy : the answer is in your two times and your two maximum task durations.
# The salt splits receiver 1 over eight partitions, so the maximum task duration should fall.
# If the total time did not fall with it, the copies of receiver 1's row and the extra work
# on the ping side ate the gain.

4g

Run the next cell. It switches AQE's skew handling back on, joins pings to receivers again, and aggregates a checksum over the three selected columns so that the shuffle keeps all of them. It reads the final plan for SortMergeJoin(skew=true). If AQE did not split a partition, the cell lowers AQE's thresholds and tries again : you would not do that at work. Open the query in the SQL tab of the Spark UI as well.

Hint 1

The slide AQE Skew Join names the two settings that decide whether AQE splits a partition.

Check your output
AQE split a partition : True

If it prints False, the cell lowers the thresholds and prints a second line. The plan that follows contains SortMergeJoin(skew=true) when a partition was split.

4h

Run the cell that puts the settings back.

Step 5Broadcast join

5a

Run the cell of Step 5. It builds carriers, a table of eight carriers typed in by hand, and sets the broadcast threshold to Spark's own 10 MB, so the step behaves the same on every platform. It joins the flights to carriers twice and averages the arrival delay per carrier name : once as it is, once with broadcast(carriers). It prints the join each plan uses.

Check your output
+------------+-------------+
|carrier_name|avg_arr_delay|
+------------+-------------+
|         ...|          ...|
+------------+-------------+

as it is     [...]
broadcast()  [...]

One row per carrier. Each list names the join strategy in the plan, and the side it builds.

5b

Call explain() on plain and on hinted. Which join does each plan use, and which side does it build, BuildLeft or BuildRight?

Hint 1

The slide Analyze Broadcast Join shows what a broadcast plan looks like.

Check your output
== Physical Plan ==
...Join ...

Two plans, one after the other.

Show solutionHide solution
python
plain.explain()
hinted.explain()
# hinted : BroadcastHashJoin with BuildRight, the right side is the carriers.
# plain  : normally SortMergeJoin, see 5c.

5c

Why did Spark not broadcast an eight-row table on its own? Write the answer in a comment.

Hint 1

The slide Broadcast Join Threshold says what Spark compares with the threshold. Where does carriers come from, and what does Spark know about its size?

Show solutionHide solution
python
# A DataFrame built from a Python list carries no size estimate, so Spark cannot know it is small.
# A table read from parquet, such as airports, is broadcast without being asked.
# Spark estimates the other side too : on a small copy of the flights, the two columns the join
# needs can fall under the threshold, and the plan then broadcasts the flights, BuildLeft.

5d

Open both queries in the Spark UI. Which number disappears from the flights side of the join when you broadcast?

Hint 1

The slides Shuffle and Broadcasting say what moves across the network in each join.

Show solutionHide solution
python
# The shuffle write of the flights side. The sort-merge join shuffles the flights to meet the
# carriers ; the broadcast join leaves them where they are and sends the small table to them.

Step 6Second-busiest carrier

For every departure airport, find the carrier with the second most departures and the carrier with the fewest.

6a

Count the departures per origin and carrier into per_carrier, with the count named departures.

Show solutionHide solution
python
per_carrier = (flights.groupBy("origin", "carrier")
                      .agg(F.count("*").alias("departures")))

6b

Rank the carriers twice over the same partitioning, once from most departures and once from fewest, with dense_rank. Add the two ranks to per_carrier as rank_desc and rank_asc, in ranked. Fill in the cell with ranked = None.

Hint 1

You need two window specifications that both partition by origin and order in opposite directions. The slides Worst Delay per Airport and Ranking Functions of module 17 show the shape.

Hint 2

F.desc("departures") orders from most to fewest.

Show solutionHide solution
python
most_first = Window.partitionBy("origin").orderBy(F.desc("departures"))
least_first = Window.partitionBy("origin").orderBy("departures")

ranked = (per_carrier
    .withColumn("rank_desc", F.dense_rank().over(most_first))
    .withColumn("rank_asc", F.dense_rank().over(least_first)))

6c

Keep rank_desc == 2 for the second-busiest carrier and rank_asc == 1 for the least busy, and build answer with one row per airport and the columns origin, carriers, second_busiest and least_busy. Handle two edge cases :

  • an airport served by one carrier has no second-busiest ; label it "only carrier"
  • a least busy carrier that is also the second-busiest counts only as second-busiest

Show the ten airports with the most carriers.

Hint 1

Count the carriers per airport from per_carrier. Join the three small tables back together with left joins on origin.

Hint 2

F.when(...).otherwise(...) handles both edge cases, one per column.

Check your output
+------+--------+--------------+----------+
|origin|carriers|second_busiest|least_busy|
+------+--------+--------------+----------+
|   ...|     ...|           ...|       ...|
+------+--------+--------------+----------+

Ten rows, carriers falling from the top. The values depend on the data.

Show solutionHide solution
python
second = ranked.filter(F.col("rank_desc") == 2).select("origin", F.col("carrier").alias("second_busiest"))
least = ranked.filter(F.col("rank_asc") == 1).select("origin", F.col("carrier").alias("least_busy"))
served = per_carrier.groupBy("origin").agg(F.count("*").alias("carriers"))

answer = (served.join(second, "origin", "left").join(least, "origin", "left")
    .withColumn("second_busiest",
                F.when(F.col("carriers") == 1, F.lit("only carrier")).otherwise(F.col("second_busiest")))
    .withColumn("least_busy",
                F.when(F.col("least_busy") == F.col("second_busiest"), F.lit(None)).otherwise(F.col("least_busy"))))
answer.orderBy(F.desc("carriers")).show(10)

6d

Look at the plan of ranked. How many Exchange nodes does it have, and how many Sort nodes? Would a window partitioned by carrier in place of origin add a shuffle? Write the answer to the second question in a comment.

Hint 1

The slide EnsureRequirements says when Spark adds an Exchange to a plan.

Check your output
Exchange nodes : ... ; Sort nodes : ...
Show solutionHide solution
python
tree = plan_text(ranked).split(chr(10) + chr(10))[0]
print("Exchange nodes :", tree.count(" Exchange ("), "; Sort nodes :", tree.count("Sort ("))

# Yes. A window needs the rows of its partition together. Both windows here partition by origin,
# so one Exchange serves both. A window partitioned by carrier needs the rows grouped another
# way, and that is a further Exchange.

Step 7Tail hashes

Add a column hashed_tail to every flight :

  • if flight_number is even, apply MD5 to tail_number once for every capital N in it, feeding each hash into the next
  • if flight_number is odd, apply SHA-256 to tail_number once

Then check whether any two flights share a hashed_tail.

7a

Write the column as a Python UDF with hashlib. A tail number or a flight number can be missing, so the function has to return None for those. Register it as a UDF returning a string and add the column to the flights as with_udf. Fill in the cell with with_udf = None.

Hint 1

The slide UDF shows a UDF with F.udf and a return type. pyspark.sql.types.StringType is the return type here.

Hint 2

tail.count("N") gives the number of capital Ns. A for loop over range(...) applies the hash that many times. The hashes come from hashlib.md5(...).hexdigest() on encoded text.

Show solutionHide solution
python
import hashlib
from pyspark.sql.types import StringType

def hash_tail(tail, flight_number):
    if tail is None or flight_number is None:
        return None
    if flight_number % 2 == 0:
        hashed = tail
        for _ in range(tail.count("N")):
            hashed = hashlib.md5(hashed.encode()).hexdigest()
        return hashed
    return hashlib.sha256(tail.encode()).hexdigest()

hash_tail_udf = F.udf(hash_tail, StringType())
with_udf = flights.withColumn("hashed_tail", hash_tail_udf("tail_number", "flight_number"))

7b

Time with_udf with measure and a noop write, which runs the UDF on every row.

Hint 1

measure("label", lambda: noop(df)) times the call and prints the seconds, and under them read, shuffle write, spill and task times.

Check your output
Python UDF                                     ... s
    read ... MB | shuffle write ... MB | spill ... MB memory, ... MB disk | ... tasks, slowest ... s, median ... s

The numbers are your own.

Show solutionHide solution
python
udf_run = measure("Python UDF", lambda: noop(with_udf))

7c

Write it without a UDF. F.sha2(col, 256) and F.md5(col) are built in ; the loop is the hard part. Count the Ns in tail_number with built-in functions first, and find the largest count in the data.

Hint 1

Take the length of the tail number, take the length of the same text with every N removed, and subtract. F.regexp_replace removes them.

Check your output
most N in one tail number : ...
Show solutionHide solution
python
without_n = F.regexp_replace("tail_number", "N", "")
n_count = F.length("tail_number") - F.length(without_n)
most_n = flights.agg(F.max(n_count)).first()[0]
print("most N in one tail number :", most_n)

7d

With the largest count known, write the loop as a chain of F.when clauses, and add the column to the flights as with_builtin. Fill in the cell with with_builtin = None. Missing values give None, as in 7a.

Hint 1

Build a list of columns : the tail number, then F.md5 of the item before it, most_n times. The clause for Ns counted k picks item k.

Hint 2

Wrap the even and the odd case in one outer F.when, the way 7a's function does with if.

Show solutionHide solution
python
chain = [F.col("tail_number")]
for _ in range(most_n):
    chain.append(F.md5(chain[-1]))
even = F.when(n_count == 0, chain[0])
for k in range(1, most_n + 1):
    even = even.when(n_count == k, chain[k])

hashed = (F.when(F.col("tail_number").isNull() | F.col("flight_number").isNull(), F.lit(None))
           .when(F.col("flight_number") % 2 == 0, even)
           .otherwise(F.sha2("tail_number", 256)))
with_builtin = flights.withColumn("hashed_tail", hashed)

7e

Time with_builtin the same way as 7b.

Check your output
built-in functions                             ... s
    read ... MB | shuffle write ... MB | ...

The numbers are your own.

Show solutionHide solution
python
builtin_run = measure("built-in functions", lambda: noop(with_builtin))

7f

Do both versions give the same column? Compare them with exceptAll on the key columns flight_date, carrier, flight_number, origin and hashed_tail. Then count the hashes that more than one flight shares.

Hint 1

select(key) on both DataFrames, then exceptAll in one direction. An empty result means every row of the first is in the second.

Check your output
rows differing : 0
hashes shared by more than one flight : ...
Show solutionHide solution
python
key = ["flight_date", "carrier", "flight_number", "origin", "hashed_tail"]
print("rows differing :", with_udf.select(key).exceptAll(with_builtin.select(key)).count())
print("hashes shared by more than one flight :",
      with_builtin.groupBy("hashed_tail").count().filter(F.col("count") > 1).count())

7g

Which version is faster, by your own times from 7b and 7e? Which plan shows BatchEvalPython? Check both plans with code.

Hint 1

The slide Without UDF says where the Python work shows up in a plan.

Check your output
UDF       BatchEvalPython in the plan : True
built-in  BatchEvalPython in the plan : False
Show solutionHide solution
python
for name, df in (("UDF", with_udf), ("built-in", with_builtin)):
    print(f"{name:9} BatchEvalPython in the plan : {'BatchEvalPython' in plan_text(df, 'simple')}")
# Faster : compare your two times. The UDF moves every row out of the JVM to Python and back.

7h

Would the built-in version still be correct if next month's data had a tail number with one more N ? Write the answer in a comment.

Show solutionHide solution
python
# No. The chain of F.when clauses has one clause for every count up to the largest count found
# in today's data. A tail number with more Ns matches none of them and gives null. The UDF
# loops as often as it needs to.

Step 8Slow query

The starter has a slow query. Use it, or write one of your own against the course data. Made up, not from work : never paste a query or data from work into this environment. It is a foreign tenant, and everything here is made up or public.

8a

Run the cell with the slow query. It times the query twice with measure and prints its formatted plan. Write down the slower and the faster run.

Check your output
slow query, first run                           ... s
    read ... MB | shuffle write ... MB | ...
slow query, second run                          ... s
    read ... MB | shuffle write ... MB | ...
== Physical Plan ==
...

The times are your own, and the second run is often faster than the first.

8b

Read the plan of slow. Mark every Exchange, the join strategy, the Python node, and what the scan filters on.

Hint 1

The slide What to Measure First lists what to look for, and the slide Partition Pruning says how a filter shows up in a scan.

Check your output
BatchEvalPython          : True
join, slow query          : [...]
Show solutionHide solution
python
text = plan_text(slow)
print("BatchEvalPython          :", "BatchEvalPython" in text)
print("join, slow query          :", join_used(slow))
# PartitionFilters lists what the scan prunes on. The filter on year(flight_date) is not among them.

8c

In the Spark UI, write down the four numbers from the slides for the slowest stage of the slow query : duration, shuffle read and write, maximum against median task time, and spill. The line under each measure result gives a total of the same kind for the whole query.

Hint 1

The slide What to Measure First names the four numbers.

8d

Change one thing. Put the changed query in faster, time it with measure in the cell provided, and read the four numbers again. Write one line : what you changed, which number moved, and by how much. A change that moved nothing is worth a line too.

Hint 1

Which node in the plan holds most of the time ? Change that one.

Hint 2

A change to the filter, the Python function or the join each has a slide : Partition Pruning, Without UDF, Broadcasting.

Check your output
1. filter on year, the partition column        ... s
    read ... MB | shuffle write ... MB | ...

The read figure falls and the time follows it. The numbers are your own.

Show solutionHide solution
python
band = (F.when(F.col("dep_delay").isNull(), "cancelled")
         .when(F.col("dep_delay") <= 0, "early or on time")
         .when(F.col("dep_delay") <= 15, "up to 15")
         .when(F.col("dep_delay") <= 60, "15 to 60")
         .otherwise("over 60"))

pruned = (flights
    .filter(F.col("year") == 2024)
    .withColumn("band", delay_band_udf("dep_delay"))
    .join(carriers, "carrier")
    .groupBy("carrier_name", "band").count())

change_1 = measure("1. filter on year, the partition column", lambda: noop(pruned))
# The filter on year cuts the bytes read : year is a partition column, year(flight_date) is not.

8e

Repeat 8d until you run out of ideas, one change at a time, each with its own line.

Check your output
2. built-in functions for the band             ... s
3. broadcast the carriers                      ... s

same answer : True
PartitionFilters on year : True
BatchEvalPython          : False
join, slow query          : [...]
join, after the changes   : ['BroadcastHashJoin', 'BuildRight']

The times are your own. Your list of changes may differ from these three.

Show solutionHide solution
python
no_udf = (flights
    .filter(F.col("year") == 2024)
    .withColumn("band", band)
    .join(carriers, "carrier")
    .groupBy("carrier_name", "band").count())
broadcast_too = (flights
    .filter(F.col("year") == 2024)
    .withColumn("band", band)
    .join(broadcast(carriers), "carrier")
    .groupBy("carrier_name", "band").count())

change_2 = measure("2. built-in functions for the band", lambda: noop(no_udf))
change_3 = measure("3. broadcast the carriers", lambda: noop(broadcast_too))

print("same answer :", sorted(slow.collect()) == sorted(broadcast_too.collect()))
text = plan_text(broadcast_too)
print("PartitionFilters on year :", "PartitionFilters: [isnotnull(year" in text)
print("BatchEvalPython          :", "BatchEvalPython" in text)
print("join, slow query          :", join_used(slow))
print("join, after the changes   :", join_used(broadcast_too))
# The built-in band removes the Python node. The hint makes the carriers the side that is copied :
# without it Spark shuffles both sides, or on a small copy of the flights may broadcast the flights.

Step 9Clean up

9a

Run the cell that unpersists anything you cached and puts back every setting you changed.

Check your output
settings back to the session's defaults

9b

Click the session indicator and stop the session. A stopped session hands its executors back to the class.

If time permits

  • Rerun the salted join from Step 4 with 2, 8 and 32 salt values. Where does the maximum task duration stop falling, and what happens to the size of the salted receivers?
  • Cache flights, then repeat the counts from Step 2 twice. Compare the first and the second run with the uncached ones. Did the cache pay for itself at all?
  • Write one month of flights to your own Files/scratch, partitioned by carrier, once as it is and once after repartition("carrier"). Count the files each write produced. The solution notebook has cells for all three.

Tried it yourself first?

The solution is a spoiler. Work through the hints first : a wrong attempt teaches more than a solution you only read.