AWS Training
Modules Listen Certification
0:00 0:00

← PySpark

PS0 Lab — build the skew, then fix it three ways

Target: any PySpark you can run — local pyspark, a notebook, Glue interactive sessions, or EMR. Parts 1–5 run locally, which is the point: you can learn all of this without a cluster.

Cost: zero locally. On Glue/EMR, a few minutes of a small worker.

Predict before you run, every time. Write the prediction down.


Setup

from pyspark.sql import SparkSession
from pyspark.sql.functions import *
from pyspark.sql.types import *

spark = (SparkSession.builder.master("local[4]").appName("ps0-lab")
         .config("spark.sql.shuffle.partitions", "8")
         .getOrCreate())

print("Spark version:", spark.version)
print("arrow enabled:", spark.conf.get("spark.sql.execution.arrow.pyspark.enabled", "unset"))

Q1. What Spark version are you on? Does the Arrow default match what lesson 3 quoted from the current docs? If it differs, that is the lesson — write down which docs you should be reading.

# A deliberately skewed dataset: 1M rows, one whale key, plus nulls
big = (spark.range(0, 1_000_000).withColumnRenamed("id", "row_id")
       .withColumn("customer_id",
           when(col("row_id") % 100 < 70, lit("WHALE"))          # 70% one key
           .when(col("row_id") % 100 < 75, lit(None))            # 5% null
           .otherwise(concat(lit("c"), (col("row_id") % 5000).cast("string"))))
       .withColumn("amount", (rand() * 100).cast("decimal(10,2)")))

dim = (spark.range(0, 5000)
       .withColumn("customer_id", concat(lit("c"), col("id").cast("string")))
       .withColumn("segment", when(col("id") % 2 == 0, "A").otherwise("B"))
       .select("customer_id", "segment")
       .union(spark.createDataFrame([("WHALE", "A")], ["customer_id", "segment"])))

Part 1 — Lazy evaluation

missing = spark.read.parquet("/tmp/does-not-exist-abc123")   # line A
filtered = missing.filter(col("x") > 1)                      # line B
result   = filtered.count()                                  # line C

Q2. Predict which line raises. Run it. Which line does the traceback name, and why?

Q3. Time these three, then explain the numbers:

import time
def t(label, fn):
    s = time.time(); r = fn(); print(label, round(time.time()-s, 2), "s"); return r

expensive = big.join(dim, "customer_id", "left").filter(col("amount") > 50)
t("count 1", lambda: expensive.count())
t("count 2", lambda: expensive.count())

Q4. Now add .cache() and repeat. What changed, and why was the first cached call not faster?

Q5. expensive.unpersist(). Why does this matter on a real cluster?


Part 2 — Boundary crossings

Q6. Scan this and list every boundary crossing, and whether it's bounded:

d = big.filter(col("amount") > 10)
d = d.withColumn("up", upper(col("customer_id")))
d = d.groupBy("customer_id").agg(sum("amount").alias("t"))
top = d.orderBy(desc("t")).limit(20).collect()
d.write.mode("overwrite").parquet("/tmp/ps0-out")

Q7. Compare three implementations of the same transformation and time each writing to disk:

py_udf = udf(lambda s: s.upper() if s else None, StringType())

@pandas_udf(StringType())
def pd_udf(s: pd.Series) -> pd.Series:
    return s.str.upper()

t("builtin",    lambda: big.withColumn("u", upper(col("customer_id"))).write.mode("overwrite").parquet("/tmp/a"))
t("pandas_udf", lambda: big.withColumn("u", pd_udf(col("customer_id"))).write.mode("overwrite").parquet("/tmp/b"))
t("python_udf", lambda: big.withColumn("u", py_udf(col("customer_id"))).write.mode("overwrite").parquet("/tmp/c"))

Record your three numbers. Q8. What's the ratio on your data? Why shouldn't you quote it as a general fact?

Q9. The pushdown demo — run both and compare explain():

big.filter(col("customer_id") == "WHALE").explain()
big.filter(py_udf(col("customer_id")) == "WHALE").explain()

What appears in one plan and not the other? What does that cost?


Part 3 — Skew ⚠️

Q10. Measure the skew before fixing anything:

big.groupBy("customer_id").count().orderBy(desc("count")).show(5)
big.groupBy(spark_partition_id().alias("pid")).count().orderBy(desc("count")).show(10)

What fraction of rows is in the largest key? The largest partition?

Q11. Run the naive join and time it. Look at the Spark UI stage detail (localhost:4040). Record max vs median task duration.

t("naive join", lambda: big.join(dim, "customer_id", "left")
                          .groupBy("segment").agg(sum("amount"))
                          .write.mode("overwrite").parquet("/tmp/j1"))

Q12. Fix 1 — filter the junk. Drop the null keys before the join. Time it. How much did it help, and why were nulls costing anything at all?

Q13. Fix 2 — broadcast. broadcast(dim). Time it. Why does skew stop mattering? What would make this a bad idea?

Q14. Fix 3 — salt. Implement salting for the WHALE key with N=10 (lesson 4 has the shape). Time it. What did it cost you that the other two didn't?

Q15. Rank the three fixes for this dataset and justify the order. Would your ranking change if dim had 50 million rows?


Part 4 — repartition vs coalesce ⚠️

def slow_map(df):
    return df.withColumn("h", sha2(col("customer_id"), 256))

t("coalesce(1)",    lambda: slow_map(big).coalesce(1).write.mode("overwrite").parquet("/tmp/c1"))
t("repartition(1)", lambda: slow_map(big).repartition(1).write.mode("overwrite").parquet("/tmp/r1"))

Q16. Predict which is faster before running. Which was it, and why? Check the stage detail — how many tasks ran the sha2 in each case?

Q17. Explain in one sentence why "coalesce doesn't shuffle so it's cheaper" is incomplete.


Part 5 — Small files and structure

Q18. Write big without repartitioning, then count the output files. Now repartition(4) and count again. What's the trade-off?

import os
print(len([f for f in os.listdir("/tmp/ps0-out") if f.endswith(".parquet")]))

Q19. Refactor Part 3's pipeline into two pure functions (DataFrame -> DataFrame) plus a thin main(spark, in_path, out_path).

Q20. Write one pytest test for one of those functions using a 4-row hand-built DataFrame. Include a null and a duplicate.

Q21. Time your test suite with spark.sql.shuffle.partitions at its default versus at 2. Record both numbers.

Q22. Add assert_unique and assert_not_empty to the job (lesson 5). Deliberately break the input so assert_unique fires. What would have happened without it?


Part 6 — On AWS (optional)

Skip unless you have Glue or EMR.

Q23. Run the Part 1 setup cell on Glue or EMR. What is spark.version? How does it compare to your local version and to the current docs?

Q24. Print all config (getConf().getAll()). Which of the Arrow settings from lesson 3 are set, and to what? Any surprises?

Q25. Put a print() in the driver and a print() inside a UDF. Find both in CloudWatch. Which log group is each in?

Q26. Write the answers to Q23–Q25 into your team's runbook. That page is the reason nobody else has to repeat this.

Teardown

spark.stop()
# rm -rf /tmp/ps0-out /tmp/a /tmp/b /tmp/c /tmp/j1 /tmp/c1 /tmp/r1

Done when you can

Facilitator notes