Pair your devices with a code and playback position follows you: pause on this device, hit resume on the other. Position is saved to the site every minute and on pause.
Open this panel on your other device and enter the same code.
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.
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"])))
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?
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?
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?
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.
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?
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.
spark.stop()
# rm -rf /tmp/ps0-out /tmp/a /tmp/b /tmp/c /tmp/j1 /tmp/c1 /tmp/r1
filter is worse than a UDF in a withColumncoalesce upstream-parallelism catchPushedFilters vanish when a UDF enters the
predicate makes the "optimiser goes blind" point permanent.coalesce(1) wins. Let them be wrong
before explaining.