SpiralTrain
Exercises › Block 2 · Exercise 4

Block 2 · Exercise 4

PySpark and SQL

Open the notebook 16-pyspark-sql-starter in your workspace and run its setup cell, which starts the Spark session. It registers the views flights and airports. Steps 2, 5 and 8 have a challenge cell. In the other steps the query is already in a cell : run it and answer the questions.

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 16-pyspark-sql-solution, under every part on the exercise site, and at the back of the exercises PDF.

Step 1Data in SQL

1a

Run the first cell under Step 1. It counts the flights.

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

The count depends on the size of the data set you loaded.

Show solutionHide solution
python
spark.sql("SELECT count(*) AS flights FROM flights").show()

1b

Run the DESCRIBE in the same cell. It reads the schema. Which took longer, the count or the DESCRIBE? Write the answer in a comment.

Hint 1

One of the two has to read the data. The other only has to look something up.

Check your output
+-------------+-----------+-------+
|col_name     |data_type  |comment|
+-------------+-----------+-------+
|flight_date  |date       |NULL   |
|carrier      |string     |NULL   |
|flight_number|int        |NULL   |
|tail_number  |string     |NULL   |
|origin       |string     |NULL   |
|dest         |string     |NULL   |
|dep_delay    |double     |NULL   |
...
+-------------+-----------+-------+

More rows follow for the remaining columns.

Show solutionHide solution
python
spark.sql("DESCRIBE flights").show(20, truncate=False)
# DESCRIBE reads the metadata and count reads the data, so one is instant and the other is not.

Step 2The four clauses

2a

Assign to by_carrier a query with SELECT, WHERE, GROUP BY and ORDER BY. It gives the flight count and the average departure delay per carrier, worst first. Put WHERE dep_delay IS NOT NULL in it. Show the result. Which carrier comes out worst?

Hint 1

The three columns are carrier, flights and avg_delay. Order by the last one, descending.

Hint 2

Round the average with round(avg(dep_delay), 2).

Check your output
+-------+-------+---------+
|carrier|flights|avg_delay|
+-------+-------+---------+
|    ...|    ...|      ...|
+-------+-------+---------+

One row per carrier, with the highest avg_delay first. The values depend on the data set you loaded.

Show solutionHide solution
python
by_carrier = spark.sql("""
    SELECT carrier,
           count(*) AS flights,
           round(avg(dep_delay), 2) AS avg_delay
    FROM flights
    WHERE dep_delay IS NOT NULL
    GROUP BY carrier
    ORDER BY avg_delay DESC
""")

by_carrier.show()

2b

Take the WHERE out of a copy of the query and run it again. Write a comment on what changes and what does not.

Hint 1

The average skips nulls on its own. Compare the flights column with the one from 2a.

Check your output
+-------+-------+---------+
|carrier|flights|avg_delay|
+-------+-------+---------+
|    ...|    ...|      ...|
+-------+-------+---------+

The same three columns as in 2a.

Show solutionHide solution
python
unfiltered = spark.sql("""
    SELECT carrier,
           count(*) AS flights,
           round(avg(dep_delay), 2) AS avg_delay
    FROM flights
    GROUP BY carrier
    ORDER BY avg_delay DESC
""")

unfiltered.show()
# The average ignores nulls whether you filter them or not, so avg_delay does not move.
# The count does : without the filter the two columns describe different sets of flights.

Step 3HAVING

3a

Run the cell under Step 3. It keeps only the carriers with more than a hundred thousand flights. Does the ranking change once the small carriers are gone, compared with 2a?

Hint 1

HAVING filters groups, not rows. The DataFrame API has no having() method.

Check your output
+-------+-------+---------+
|carrier|flights|avg_delay|
+-------+-------+---------+
|    ...|    ...|      ...|
+-------+-------+---------+

Only the carriers with more than 100000 flights. The values depend on the data set you loaded.

Show solutionHide solution
python
spark.sql("""
    SELECT carrier, count(*) AS flights, round(avg(dep_delay), 2) AS avg_delay
    FROM flights
    GROUP BY carrier
    HAVING count(*) > 100000
    ORDER BY avg_delay DESC
""").show()
# HAVING filters groups where WHERE filters rows.
# In the DataFrame API a filter after groupBy().agg() does the same.

Step 4Views

4a

Run the cell under Step 4. It registers a view very_late of the flights delayed more than an hour, checks that the view exists, counts through it and drops it. Which part took the time, registering the view or counting through it? Write the answer in a comment.

Hint 1

Registering gives a name to a query. Counting has to run one.

Check your output
exists : True
+----+
|rows|
+----+
| ...|
+----+

exists after the drop : False
Show solutionHide solution
python
spark.sql("""
    SELECT flight_date, carrier, origin, dest, dep_delay
    FROM flights WHERE dep_delay > 60
""").createOrReplaceTempView("very_late")

print("exists :", spark.catalog.tableExists("very_late"))
spark.sql("SELECT count(*) AS rows FROM very_late").show()
spark.catalog.dropTempView("very_late")
print("exists after the drop :", spark.catalog.tableExists("very_late"))
# Registering a view reads nothing. It names a query plan, and the plan runs every time
# the name is queried. A view is not a cache.

Step 5Subquery

5a

Assign to worse a query for the carriers that are worse than the average carrier. Show it. The overall average of dep_delay goes in a subquery inside the HAVING clause.

Hint 1

HAVING avg(dep_delay) > (...), with a SELECT avg(dep_delay) FROM flights between the parentheses.

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

carriers above the overall average : ...

The carriers whose avg_delay is above the overall average, the worst first.

Show solutionHide solution
python
worse = spark.sql("""
    SELECT carrier, round(avg(dep_delay), 2) AS avg_delay
    FROM flights
    GROUP BY carrier
    HAVING avg(dep_delay) > (SELECT avg(dep_delay) FROM flights)
    ORDER BY avg_delay DESC
""")

worse.show()
print("carriers above the overall average :", worse.count())

5b

Compute the overall average of dep_delay and the average of the per-carrier averages, each with its own query. Compare the two numbers. Write a comment on which of the two you would put in a report, and why.

Hint 1

The second one needs a subquery in the FROM clause that returns one average per carrier.

Check your output
+---------------+
|overall_average|
+---------------+
|            ...|
+---------------+

+-------------------+
|average_of_averages|
+-------------------+
|                ...|
+-------------------+

The two numbers are different.

Show solutionHide solution
python
spark.sql("""
    SELECT round(avg(dep_delay), 3) AS overall_average FROM flights
""").show()
spark.sql("""
    SELECT round(avg(a), 3) AS average_of_averages FROM
    (SELECT avg(dep_delay) AS a FROM flights GROUP BY carrier)
""").show()
# The overall average weights every flight equally. The average of averages weights every
# carrier equally. Neither is wrong ; knowing which one you computed is the job.

Step 6Join

6a

Run the cell under Step 6. It joins the airports on the origin and reports the average delay per departure city, for the cities with more than twenty thousand flights.

Check your output
+----+-------+---------+
|city|flights|avg_delay|
+----+-------+---------+
| ...|    ...|      ...|
+----+-------+---------+

Up to ten cities, the highest avg_delay first.

Show solutionHide solution
python
by_city = spark.sql("""
    SELECT a.municipality AS city,
           count(*) AS flights,
           round(avg(f.dep_delay), 2) AS avg_delay
    FROM flights f
    JOIN airports a ON f.origin = a.iata_code
    GROUP BY a.municipality
    HAVING count(*) > 20000
    ORDER BY avg_delay DESC
""")

by_city.show(10, truncate=False)

6b

Assign the query to by_city and call .explain() on it. Does the plan say BroadcastHashJoin or SortMergeJoin? Then open the SQL tab in the Spark UI and find the query. How many stages did it take?

Hint 1

Look for the join in the == Physical Plan ==, from the bottom up. Check how big the airports table is next to the flights.

Check your output
== Physical Plan ==
AdaptiveSparkPlan isFinalPlan=false
+- Sort ...
   +- Exchange rangepartitioning(...)
      +- Filter (count(1) > 20000)
         +- HashAggregate(...)
            ...
               +- (the join node, with the airports table on one side)

The exact nodes, numbers and order of the lines vary with the Spark version and the data.

Show solutionHide solution
python
by_city.explain()
# The airports table is small, so Spark broadcasts it : BroadcastHashJoin.

Step 7Window in SQL

7a

Run the cell under Step 7. It ranks the carriers by average delay inside each origin airport and keeps the worst carrier per airport. Which airport and carrier pair is worst overall? Why does this need a window rather than a GROUP BY? Write the second answer in a comment.

Hint 1

Look at what the query returns for each airport : an aggregate, or a whole row?

Check your output
+------+-------+-------+---------+
|origin|carrier|flights|avg_delay|
+------+-------+-------+---------+
|   ...|    ...|    ...|      ...|
+------+-------+-------+---------+

Up to ten rows, one per airport, the highest avg_delay first.

Show solutionHide solution
python
spark.sql("""
    WITH per_pair AS (
        SELECT origin, carrier,
               count(*) AS flights,
               round(avg(dep_delay), 2) AS avg_delay
        FROM flights
        WHERE dep_delay IS NOT NULL
        GROUP BY origin, carrier
        HAVING count(*) > 5000
    ),
    ranked AS (
        SELECT per_pair.*,
               rank() OVER (PARTITION BY origin ORDER BY avg_delay DESC) AS worst_rank
        FROM per_pair
    )
    SELECT origin, carrier, flights, avg_delay
    FROM ranked WHERE worst_rank = 1
    ORDER BY avg_delay DESC
    LIMIT 10
""").show()
# A GROUP BY collapses the rows, and the answer is a whole row : the worst carrier at each airport.

Step 8Same query twice

8a

Assign to api the query of 2a as a DataFrame chain, with filter, groupBy, agg and orderBy.

Hint 1

F.count("*") and F.round(F.avg("dep_delay"), 2) go inside agg, each with an alias.

Show solutionHide solution
python
api = (flights
    .filter(F.col("dep_delay").isNotNull())
    .groupBy("carrier")
    .agg(F.count("*").alias("flights"),
         F.round(F.avg("dep_delay"), 2).alias("avg_delay"))
    .orderBy(F.desc("avg_delay")))

8b

Compare api with by_carrier from 2a in three ways, and print a line for each :

  • Do they return the same rows?
  • Do the plan strings match as text?
  • Do the plans match once the counters Spark puts in the plan text are stripped out?

df._jdf.queryExecution().executedPlan().toString() gives the plan of a DataFrame as a string. Also call explain() on each and read the plans side by side.

Hint 1

Look at the plan text for count#132L. The number after # is a counter that differs from one plan to the next.

Hint 2

re.sub(r"#\d+L?", "#x", plan) replaces the counters. plan_id= numbers need the same treatment.

Check your output
same answer : True
raw text equal : False
same plan      : True

The two explain() outputs come before these lines and show the same nodes with different numbers.

Show solutionHide solution
python
import re

api.explain()
by_carrier.explain()

def plan_of(df):
    return df._jdf.queryExecution().executedPlan().toString()

def normalised(plan):
    return re.sub(r"plan_id=\d+", "plan_id=x", re.sub(r"#\d+L?", "#x", plan))

print("same answer :", api.collect() == by_carrier.collect())
print("raw text equal :", plan_of(api) == plan_of(by_carrier))
print("same plan      :", normalised(plan_of(api)) == normalised(plan_of(by_carrier)))
# Spark numbers every expression and plan node as it builds them, so count#132L in one plan
# is count#135L in the other. The plans describe the same work : there is no performance
# argument either way.

Step 9Save a table

9a

Run the cell under Step 9. It writes by_carrier as a table, queries it by name, lists the catalog and drops the table. What does tableType say for your table, and for the views? What would happen to each if you stopped the session now?

The table lands in the work lakehouse of your own workspace, and nobody else sees it. In the shared workspace, PySpark-Shared, everybody sees a table, so put your account name in its name there.

Hint 1

The catalog lists views and tables together. Read the second column.

Check your output
+-------+-------+---------+
|carrier|flights|avg_delay|
+-------+-------+---------+
|    ...|    ...|      ...|
+-------+-------+---------+

my_carrier_delays      ...
flights                ...
airports               ...

dropped : True

The second column is the tableType. The order of the catalog lines may differ.

Show solutionHide solution
python
table_name = "my_carrier_delays"

by_carrier.write.mode("overwrite").saveAsTable(table_name)
spark.sql(f"SELECT * FROM {table_name} ORDER BY avg_delay DESC LIMIT 5").show()

for t in spark.catalog.listTables():
    print(f"{t.name:22} {t.tableType}")

spark.sql(f"DROP TABLE IF EXISTS {table_name}")
print("dropped :", not spark.catalog.tableExists(table_name))
# A view is TEMPORARY and a saved table is MANAGED. The views go with the session ;
# only the table is still there for the next one.

Step 10Stop the session

10a

Click the session indicator and stop the session.

If time permits

  • Ask a question of your own about the flights, in SQL, and answer it. Bring a real one from your own work if the shape fits.
  • Rewrite step 7 without the window function, joining back to a grouped subquery instead. Which is easier to read, and which would you rather maintain?
  • Add date_format(flight_date, 'EEEE') to step 2 and find the worst carrier and weekday pair. Does any carrier have a bad day of the week?

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.