Menu

PySpark course · Lesson 4 of 8

PySpark Aggregations, Pivot and Temporary Views

Aggregate PySpark DataFrames with groupBy and agg, build subtotals with rollup and cube, pivot and unpivot data, and share DataFrames with SQL through temp views.

  • Beginner
  • 16 min read
  • Updated Oct 2026
On this page
  1. Sample data
  2. Grouping and aggregating with groupBy
  3. What it is
  4. A worked example
  5. Useful aggregate functions
  6. Filtering after aggregation (HAVING)
  7. Empty input
  8. Subtotals: rollup and cube
  9. How it runs
  10. Pitfalls
  11. In interviews
  12. Pivoting and unpivoting
  13. What it is
  14. Always pass the value list
  15. Unpivot
  16. Pitfalls
  17. In interviews
  18. Temporary views and global temporary views
  19. What they are
  20. Scope: session vs application
  21. Views are re-evaluated, and names are resolved late
  22. Pitfalls
  23. In interviews
  24. Practice questions
  25. Key takeaways

Most reports, metrics and data-quality checks are aggregations: totals per day, counts per customer, averages per region. In PySpark, an aggregation is also where the first shuffle usually happens, so understanding how groupBy runs matters for cost as well as correctness. This lesson covers grouping and aggregate functions, pivoting between long and wide shapes, and temporary views, which let you move between DataFrame code and SQL in the same job.

Sample data

A small sales table with the awkward cases built in: a missing amount, a missing region and repeated reps.

from pyspark.sql import SparkSession, functions as F

spark = SparkSession.builder.master("local[2]").appName("aggregations").getOrCreate()
spark.sparkContext.setLogLevel("ERROR")
spark.conf.set("spark.sql.shuffle.partitions", "4")   # small data: avoid 200 tiny tasks

sales = spark.createDataFrame(
    [("2026-01", "north", "laptop", "asha", 1200.0),
     ("2026-01", "north", "phone",  "ben",   600.0),
     ("2026-01", "south", "laptop", "chen",  1100.0),
     ("2026-01", "south", "phone",  "chen",  None),
     ("2026-02", "north", "laptop", "asha",  1300.0),
     ("2026-02", "north", "tablet", "ben",   400.0),
     ("2026-02", "south", "phone",  "dara",  650.0),
     ("2026-02", None,    "phone",  "erin",  500.0)],
    "month STRING, region STRING, product STRING, rep STRING, amount DOUBLE",
)
sales.show()
+-------+------+-------+----+------+
|  month|region|product| rep|amount|
+-------+------+-------+----+------+
|2026-01| north| laptop|asha|1200.0|
|2026-01| north|  phone| ben| 600.0|
|2026-01| south| laptop|chen|1100.0|
|2026-01| south|  phone|chen|  NULL|
|2026-02| north| laptop|asha|1300.0|
|2026-02| north| tablet| ben| 400.0|
|2026-02| south|  phone|dara| 650.0|
|2026-02|  NULL|  phone|erin| 500.0|
+-------+------+-------+----+------+

Grouping and aggregating with groupBy

What it is

groupBy(cols) collects rows that share the same values in the grouping columns; agg(...) then reduces each group to one row using aggregate functions. It is the DataFrame form of SQL’s GROUP BY. The output has one row per distinct key combination, with the grouping columns followed by the aggregates.

A worked example

by_region = (
    sales.groupBy("region")
         .agg(F.count("*").alias("rows"),
              F.count("amount").alias("priced_rows"),
              F.sum("amount").alias("revenue"),
              F.round(F.avg("amount"), 1).alias("avg_amount"),
              F.max("amount").alias("biggest"),
              F.countDistinct("rep").alias("reps"))
         .orderBy(F.col("region").asc_nulls_last())
)
by_region.show()
+------+----+-----------+-------+----------+-------+----+
|region|rows|priced_rows|revenue|avg_amount|biggest|reps|
+------+----+-----------+-------+----------+-------+----+
| north|   4|          4| 3500.0|     875.0| 1300.0|   2|
| south|   3|          2| 1750.0|     875.0| 1100.0|   2|
|  NULL|   1|          1|  500.0|     500.0|  500.0|   1|
+------+----+-----------+-------+----------+-------+----+

Three NULL rules are visible here:

  • NULL keys form their own group. The row with no region is not dropped; it becomes the NULL group.
  • count("*") counts rows; count(column) counts non-NULL values. South has 3 rows but only 2 priced rows.
  • sum, avg, max ignore NULLs. South’s average is 1750 / 2 = 875, not 1750 / 3.

Shortcuts exist for single aggregates (groupBy("region").sum("amount"), .count()), but they produce generated column names like sum(amount), so agg with alias is clearer in pipelines:

sales.groupBy("region").sum("amount").orderBy("region").show()
sales.groupBy("month", "region").count().orderBy("month", "region").show()
+------+-----------+
|region|sum(amount)|
+------+-----------+
|  NULL|      500.0|
| north|     3500.0|
| south|     1750.0|
+------+-----------+

+-------+------+-----+
|  month|region|count|
+-------+------+-----+
|2026-01| north|    2|
|2026-01| south|    2|
|2026-02|  NULL|    1|
|2026-02| north|    2|
|2026-02| south|    1|
+-------+------+-----+

Useful aggregate functions

sales.agg(F.sum("amount").alias("total"), F.count("*").alias("rows")).show()
sales.groupBy("region").agg(
    F.collect_list("rep").alias("reps_list"),
    F.collect_set("product").alias("products"),
    F.sum(F.when(F.col("product") == "laptop", F.col("amount"))).alias("laptop_revenue"),
    F.count_if(F.col("amount") >= 1000).alias("big_sales"),
    F.max_by("rep", "amount").alias("top_rep"),
).orderBy("region").show(truncate=False)
+------+----+
| total|rows|
+------+----+
|5750.0|   8|
+------+----+

+------+----------------------+-----------------------+--------------+---------+-------+
|region|reps_list             |products               |laptop_revenue|big_sales|top_rep|
+------+----------------------+-----------------------+--------------+---------+-------+
|NULL  |[erin]                |[phone]                |NULL          |0        |erin   |
|north |[asha, ben, asha, ben]|[tablet, laptop, phone]|2500.0        |2        |asha   |
|south |[chen, chen, dara]    |[laptop, phone]        |1100.0        |1        |chen   |
+------+----------------------+-----------------------+--------------+---------+-------+
  • df.agg(...) with no groupBy aggregates the whole DataFrame into one row.
  • collect_list keeps duplicates, collect_set removes them. Neither guarantees order, and both build the whole list in memory for each group, so avoid them on large groups.
  • Conditional aggregation (sum(when(...))) is the portable way to compute several filtered totals in one pass; when without otherwise returns NULL, which sum ignores.
  • max_by(x, y) returns the x from the row with the largest y: “which rep made the biggest sale” without a join or window.

Filtering after aggregation (HAVING)

There is no separate having method: filter the aggregated DataFrame.

(sales.groupBy("region").agg(F.sum("amount").alias("revenue"))
      .filter(F.col("revenue") > 2000)
      .show())
+------+-------+
|region|revenue|
+------+-------+
| north| 3500.0|
+------+-------+

Empty input

A global aggregate over no rows returns one row (with sum as NULL and count as 0); a grouped aggregate over no rows returns no rows. Downstream code that expects “one row per day” must handle both.

spark.range(0).agg(F.sum("id").alias("s"), F.count("*").alias("c")).show()
spark.range(0).groupBy("id").count().show()
+----+---+
|   s|  c|
+----+---+
|NULL|  0|
+----+---+

+---+-----+
| id|count|
+---+-----+
+---+-----+

Subtotals: rollup and cube

rollup("region", "month") returns the normal groups plus a subtotal per region and a grand total. cube returns subtotals for every combination of the columns. Subtotal rows show NULL in the rolled-up column, which collides with genuine NULL keys. grouping_id() tells them apart: 0 is a normal group, and higher values say which columns were rolled up.

(sales.rollup("region", "month")
      .agg(F.sum("amount").alias("revenue"), F.grouping_id().alias("gid"))
      .orderBy(F.col("region").asc_nulls_last(), F.col("month").asc_nulls_last())
      .show())
+------+-------+-------+---+
|region|  month|revenue|gid|
+------+-------+-------+---+
| north|2026-01| 1800.0|  0|
| north|2026-02| 1700.0|  0|
| north|   NULL| 3500.0|  1|
| south|2026-01| 1100.0|  0|
| south|2026-02|  650.0|  0|
| south|   NULL| 1750.0|  1|
|  NULL|2026-02|  500.0|  0|
|  NULL|   NULL| 5750.0|  3|
|  NULL|   NULL|  500.0|  1|
+------+-------+-------+---+

The last two rows both look like NULL, NULL. Only gid shows that 5750 is the grand total (both columns rolled up, binary 11 = 3) and 500 is the subtotal for the real NULL region (month rolled up, binary 01 = 1).

How it runs

by_region.explain()
== Physical Plan ==
AdaptiveSparkPlan isFinalPlan=false
+- Sort [region#1 ASC NULLS LAST], true, 0
   +- Exchange rangepartitioning(region#1 ASC NULLS LAST, 4), ENSURE_REQUIREMENTS, [plan_id=631]
      +- HashAggregate(keys=[region#1], functions=[count(1), count(amount#4), sum(amount#4), avg(amount#4), max(amount#4), count(distinct rep#3)])
         +- Exchange hashpartitioning(region#1, 4), ENSURE_REQUIREMENTS, [plan_id=628]
            +- HashAggregate(keys=[region#1], functions=[merge_count(1), merge_count(amount#4), merge_sum(amount#4), merge_avg(amount#4), merge_max(amount#4), partial_count(distinct rep#3)])
               +- HashAggregate(keys=[region#1, rep#3], functions=[merge_count(1), merge_count(amount#4), merge_sum(amount#4), merge_avg(amount#4), merge_max(amount#4)])
                  +- Exchange hashpartitioning(region#1, rep#3, 4), ENSURE_REQUIREMENTS, [plan_id=624]
                     +- HashAggregate(keys=[region#1, rep#3], functions=[partial_count(1), partial_count(amount#4), partial_sum(amount#4), partial_avg(amount#4), partial_max(amount#4)])
                        +- Project [region#1, rep#3, amount#4]
                           +- Scan ExistingRDD[month#0,region#1,product#2,rep#3,amount#4]

Read it bottom-up. Spark first computes partial aggregates inside each input partition (partial_sum, partial_count), then shuffles only those small partial results by key (Exchange hashpartitioning), then merges them (merge_sum). This map-side combining is why groupBy().agg() is far cheaper than collecting all rows per key. The countDistinct adds an extra shuffle on (region, rep), which is why exact distinct counts are expensive; approx_count_distinct avoids it when an estimate is acceptable.

Pitfalls

  • Grouping by a skewed key (one customer with most of the rows) makes one task do most of the work. See partitions, shuffles and skew.
  • Non-deterministic first() / last(). Without an ordering they return an arbitrary row from the group. Use max_by/min_by or a window with row_number.
  • avg of an integer column returns a double, and sum of integers returns a bigint; under ANSI mode an overflowing sum raises an error.
  • Too many shuffle partitions for small data (default 200) create many tiny tasks; AQE coalesces them at run time, but small tests run faster with a lower setting.

In interviews

Expect “what is the difference between count(*) and count(col)?”, “why is groupByKey slower than reduceByKey?” (the DataFrame groupBy().agg() does the map-side combine for you) and “write the query for revenue per region over 2000”. Mention NULL handling and the partial-then-final aggregation with a shuffle in between.

Pivoting and unpivoting

What it is

A pivot turns distinct values of one column into separate columns: one row per region with a column per month. It is the “long to wide” reshape used for reports and feature tables. Unpivot (also called melt) does the reverse.

In PySpark, pivot sits between groupBy and agg: groupBy(row keys).pivot(column whose values become columns).agg(value).

pivoted = sales.groupBy("region").pivot("month").agg(F.sum("amount"))
pivoted.orderBy("region").show()
+------+-------+-------+
|region|2026-01|2026-02|
+------+-------+-------+
|  NULL|   NULL|  500.0|
| north| 1800.0| 1700.0|
| south| 1100.0|  650.0|
+------+-------+-------+

Always pass the value list

Without a list of values, Spark runs an extra job to find the distinct values of month first, and the output columns change whenever new values appear in the data, which breaks downstream schemas. Passing the list skips that job and fixes the schema. With several aggregates, column names combine the value and the alias:

sales.groupBy("region").pivot("month", ["2026-01", "2026-02"]).agg(
    F.sum("amount").alias("rev"), F.count(F.lit(1)).alias("n")
).orderBy("region").show()
+------+-----------+---------+-----------+---------+
|region|2026-01_rev|2026-01_n|2026-02_rev|2026-02_n|
+------+-----------+---------+-----------+---------+
|  NULL|       NULL|     NULL|      500.0|        1|
| north|     1800.0|        2|     1700.0|        2|
| south|     1100.0|        2|      650.0|        1|
+------+-----------+---------+-----------+---------+

Note that even count gives NULL, not 0, for a cell with no rows: pivot cells with no matching data are NULL whatever the aggregate. (Also, count("*") is not allowed inside a pivot; count a literal instead.) Wrap the aggregate when you want zeros:

(sales.groupBy("product")
      .pivot("month", ["2026-01", "2026-02"])
      .agg(F.coalesce(F.sum("amount"), F.lit(0.0)))
      .orderBy("product").show())
+-------+-------+-------+
|product|2026-01|2026-02|
+-------+-------+-------+
| laptop| 2300.0| 1300.0|
|  phone|  600.0| 1150.0|
| tablet|   NULL|  400.0|
+-------+-------+-------+

tablet in January is still NULL: there were no rows at all for that cell, so the aggregate never ran. coalesce inside agg only fixes groups that exist but sum to NULL. To fill empty cells, apply na.fill(0) to the pivoted result.

Unpivot

DataFrame.unpivot (alias melt, Spark 3.4+) takes the identifier columns, the columns to turn into rows, and names for the new key and value columns:

wide = pivoted.filter(F.col("region").isNotNull())
long = wide.unpivot("region", ["2026-01", "2026-02"], "month", "revenue")
long.orderBy("region", "month").show()
+------+-------+-------+
|region|  month|revenue|
+------+-------+-------+
| north|2026-01| 1800.0|
| north|2026-02| 1700.0|
| south|2026-01| 1100.0|
| south|2026-02|  650.0|
+------+-------+-------+

On older Spark versions the same reshape is written with stack in selectExpr.

Pitfalls

  • Pivoting a high-cardinality column (customer ID, timestamp) produces thousands of columns. Spark limits a pivot without a value list to spark.sql.pivotMaxValues distinct values (10,000 by default) and very wide DataFrames are slow to plan. Pivot only bounded, known sets.
  • Column names from data may contain characters such as - or spaces that need backticks in SQL later. Rename them with aliases in the value list (the SQL form below) or withColumnRenamed.
  • Every pivot is an aggregation: if a cell can have several rows, decide whether sum, max or first is right.

In interviews

Common prompt: “turn monthly rows into one row per customer with a column per month” or the reverse. Mention passing explicit values, NULL in empty cells and unpivot/stack for the reverse.

Temporary views and global temporary views

What they are

A temporary view gives a DataFrame a name that Spark SQL can query. It stores no data: the name points to the DataFrame’s logical plan, which runs when queried. It lets you write part of a job in SQL and the rest in Python, and lets analysts reuse SQL they already have.

sales.createOrReplaceTempView("sales")
spark.sql("""
    SELECT region, month, SUM(amount) AS revenue
    FROM sales
    WHERE region IS NOT NULL
    GROUP BY region, month
    ORDER BY region, month
""").show()
print([t.name for t in spark.catalog.listTables()])
+------+-------+-------+
|region|  month|revenue|
+------+-------+-------+
| north|2026-01| 1800.0|
| north|2026-02| 1700.0|
| south|2026-01| 1100.0|
| south|2026-02|  650.0|
+------+-------+-------+

['sales']

SQL and DataFrame code compile to the same plans through the same optimiser, so there is no performance difference between them. Spark SQL also has a PIVOT clause, which lets you alias the generated columns:

spark.sql("""
    SELECT * FROM (SELECT region, month, amount FROM sales)
    PIVOT (SUM(amount) FOR month IN ('2026-01' AS jan, '2026-02' AS feb))
    ORDER BY region NULLS LAST
""").show()
+------+------+------+
|region|   jan|   feb|
+------+------+------+
| north|1800.0|1700.0|
| south|1100.0| 650.0|
|  NULL|  NULL| 500.0|
+------+------+------+

Scope: session vs application

Temporary view Global temporary view
Created with createOrReplaceTempView / CREATE TEMP VIEW createOrReplaceGlobalTempView / CREATE GLOBAL TEMP VIEW
Visible to The SparkSession that created it All SparkSessions in the same Spark application
Referenced as name global_temp.name (the system database global_temp)
Lives until The session ends or you drop it The application stops or you drop it
Stored data None, just a plan None, just a plan

Neither is a permanent table: nothing is written to a metastore, and another application (another job, another notebook cluster) cannot see them. For that, save a table with saveAsTable or write to storage.

sales.createOrReplaceGlobalTempView("sales_g")
other = spark.newSession()
print(other.sql("SELECT COUNT(*) AS n FROM global_temp.sales_g").collect())
try:
    other.sql("SELECT COUNT(*) FROM sales").collect()
except Exception as e:
    print(type(e).__name__, str(e).split(".")[0][:80])
[Row(n=8)]
AnalysisException [TABLE_OR_VIEW_NOT_FOUND] The table or view `sales` cannot be found

The second session sees the global view but not the session-scoped sales view.

Views are re-evaluated, and names are resolved late

A view runs its plan every time it is queried: if the source files change, a later query sees new data. A view defined in SQL that references another view also looks that name up again, so replacing the underlying view changes the dependent one:

spark.sql("CREATE OR REPLACE TEMP VIEW big_sales AS SELECT * FROM sales WHERE amount >= 1000")
print(spark.table("big_sales").count())
sales2 = sales.filter("amount < 1000")
sales2.createOrReplaceTempView("sales")
print(spark.table("sales").count(), spark.table("big_sales").count())
print(spark.catalog.dropTempView("big_sales"), spark.catalog.dropGlobalTempView("sales_g"))
3
4 0
True True

big_sales returned 3 rows, then 0 after sales was replaced with a DataFrame containing only small amounts, even though big_sales itself was never touched. In long notebooks this causes confusing results; use distinct names for each stage.

Pitfalls

  • A view is not a cache. Querying a view of an expensive plan five times runs the plan five times. Call cache() on the DataFrame (or CACHE TABLE) if you reuse it.
  • createTempView (without “OrReplace”) fails if the name exists; the “OrReplace” variants silently overwrite, which hides name collisions.
  • Global temp views require the global_temp. prefix; forgetting it gives “table not found”.
  • Spark Connect and notebooks: each client session has its own temp views, so a view made in one notebook is not visible in another attached to the same cluster unless it is global (and even then, only within the same Spark application).

In interviews

“Temp view vs global temp view vs table?” A strong answer covers scope (session, application, metastore), that views hold no data and are recomputed, and the global_temp database. A follow-up is often “does SQL run slower than DataFrame code?” (no: same optimiser and plans).

Practice questions

What is the difference between count(*), count(col) and countDistinct(col)?

count(*) counts rows, including rows where every column is NULL. count(col) counts rows where col is not NULL. countDistinct(col) counts distinct non-NULL values and needs an extra shuffle, so it is more expensive; approx_count_distinct is the cheaper estimate.

Why does groupBy().agg(sum) shuffle less data than collecting all values per key and summing them yourself?

Spark computes partial aggregates inside each partition before the shuffle, so it moves one partial sum per key per partition instead of every row. Collecting all values (collect_list and summing, or groupByKey on RDDs) shuffles every row and builds whole groups in memory.

A rollup output has two rows with region NULL and month NULL. How do you know which is the grand total?

Use grouping_id() (or grouping(col)) in the aggregation. A subtotal row has the rolled-up bits set; the grand total has all bits set. A row whose NULL came from the data has the bit for that column cleared.

Why should you pass the list of values to pivot?

Without it, Spark runs an extra job to compute the distinct values, and the output schema depends on the data, so a new value adds a column and breaks downstream consumers. With the list, there is no extra job and the schema is fixed.

You pivot with count and expect 0 for empty cells, but get NULL. Why, and how do you fix it?

A pivot cell with no matching rows has no group to aggregate, so it is NULL regardless of the function. Apply na.fill(0) (or coalesce) to the pivoted columns afterwards.

What is the difference between a temporary view, a global temporary view and a table?

A temporary view is a named plan visible only to its SparkSession. A global temporary view is visible to all sessions in the same application under global_temp. Both store no data and disappear when the session or application ends. A table (via saveAsTable or CREATE TABLE) is registered in the catalog and its data persists in storage, visible to other applications.

Key takeaways

  • groupBy().agg() reduces groups with partial aggregation, a shuffle by key and a final merge; NULL keys form their own group and aggregates ignore NULL values.
  • Use alias on every aggregate, conditional aggregation for filtered totals, and max_by/min_by instead of first for “the row with the largest value”.
  • Global aggregates on empty input return one row; grouped aggregates return none.
  • rollup and cube add subtotal rows; distinguish them from genuine NULLs with grouping_id.
  • Pivot with an explicit value list; empty cells are NULL. Use unpivot to go back to long format.
  • Temp views are session-scoped names for plans, global temp views are application-scoped under global_temp, and neither stores data.

By Data Career Hub Editorial · Last reviewed Oct 2026 · All examples run on PySpark 4.2.0 in local mode (local[2]) with spark.sql.shuffle.partitions set to 4 so the small examples run quickly. DataFrame.unpivot needs Spark 3.4 or later; count_if and max_by need Spark 3.x or later.

Progress is saved in this browser only. No account needed.

Search
Filter by type