Skip to content
Great engineers know 4 min read

How aggregations execute

Partial and final aggregation, hash vs sort aggregation, and why COUNT(DISTINCT) costs more.

TL;DR Spark aggregates in two phases: each task pre-aggregates its own rows (partial_sum), the partial results are shuffled by the grouping key, and a final aggregation merges them. The operator is a fast HashAggregate when it can be, ObjectHashAggregate for functions like collect_list, and a slower SortAggregate otherwise. Distinct aggregates add extra aggregation steps.
Read first: Narrow vs wide transformations · Catalyst and physical plans (skip if you know them)

Partial and final aggregation

SELECT country, SUM(amount) FROM orders GROUP BY country

HashAggregate(keys=[country], functions=[sum(amount)])              <- final: merge partial sums
+- Exchange hashpartitioning(country, 200)                           <- shuffle partial results
   +- HashAggregate(keys=[country], functions=[partial_sum(amount)]) <- partial: per task
      +- FileScan parquet orders [country,amount]

If a tasktask: The work for one partition in one stage, run on one CPU core. Learn more → holds a million orders from 50 countries, the partial aggregation sends 50 rows to the shuffleshuffle: Moving rows between machines so that all rows with the same key end up together. Needed by joins, groupBy and sorting, and usually the most expensive step of a job. Learn more → instead of a million. This is the same idea as a MapReduce combiner, and it is why a groupBy on a hot key is usually far less painful than a join on one: the hot key is collapsed to one row per task before the shuffle.

Partial aggregation only helps when each task sees repeated keys. Grouping by a nearly unique column (an order id) does the hashing work twice and shrinks nothing.

The three aggregation operators

OperatorChosen whenBehaviour
HashAggregateAll aggregation buffers are fixed-width mutable types: sum, count, avg, min/max of numbers and datesKeeps an off-heap hash map of key → buffer in Tungsten binary format. Fastest. If the map cannot grow, it sorts and spillsspill: Writing data to local disk because it does not fit in memory. Slower, but the job keeps running. Learn more → what it has and finishes the rest with sort-based aggregation.
ObjectHashAggregateFunctions with object buffers: collect_list, collect_set, percentile_approx, typed Dataset aggregatorsA hash map of JVM objects. After spark.sql.objectHashAggregate.sortBased.fallbackThreshold (128) keys in a task, it switches to sort-based aggregation.
SortAggregateNeither of the above applies, for example min/max of a string column in many versionsSorts rows by the grouping key, then aggregates each run of equal keys. An extra sort of every input row.
Tip: a SortAggregate on a large input is worth a second look. Sometimes a cheap rewrite moves it back to hash aggregation, for example max_by/min_by on a numeric column, or aggregating an id and joining the string back afterwards.

Distinct aggregates

A COUNT(DISTINCT user_id) per country cannot be pre-aggregated into a single number, because the same user can appear in many tasks. Spark rewrites it into two levels of aggregation:

  1. 1
    Aggregate by (country, user_id), partial then final, with a shuffle on both columns. This deduplicates users.
  2. 2
    Aggregate the result by country, counting rows, with a second shuffle.

With several distinct aggregates on different columns, such as COUNT(DISTINCT user_id) and COUNT(DISTINCT product_id) together, Spark adds an Expand operator that copies every input row once per distinct group before the first shuffle, so the data shuffled grows with the number of distinct columns.

PySparkSpark SQL · Exact vs approximate distinct counts
orders.groupBy("country").agg(F.countDistinct("user_id").alias("users"))
orders.groupBy("country").agg(F.approx_count_distinct("user_id", 0.02).alias("users"))
SELECT country, COUNT(DISTINCT user_id) AS users FROM orders GROUP BY country;
SELECT country, approx_count_distinct(user_id, 0.02) AS users FROM orders GROUP BY country;

approx_count_distinct uses HyperLogLog++ sketches, which merge like ordinary partial aggregates: one shuffle, fixed memory per group, and a relative standard error you choose (5% by default).

Spotting trouble in a plan

  • SortAggregate on a large input: an extra sort; see if a rewrite gets a HashAggregate.
  • Expand under an aggregation: several distinct aggregates multiplying rows.
  • Spill in the aggregation stagestage: A group of steps Spark can run without moving data between machines. A new stage starts at every shuffle. Learn more →: too many keys per task for memory; more shuffle partitionspartition: A chunk of a DataFrame's rows. Spark processes each partition as one task, so partitions decide how much work runs in parallel. Learn more → help.
  • collect_list on a hot key: object hash aggregation cannot shrink it; one huge array lands in one task.

Common mistakes

  • Assuming every groupBy is equally cheap — Operator choice, distinct aggregates and object buffers change the cost a lot.
  • Many COUNT(DISTINCT) on different columns in one query — The Expand operator multiplies the shuffled rows.
  • Using exact distinct counts for dashboards — approx_count_distinct is one shuffle and fixed memory, with a chosen error.
  • Salting a groupBy with only sum and count — Partial aggregation already handles hot keys for these; salting helps joins and non-decomposable aggregates more.

What you learned

  • Why a groupBy runs as a partial and a final aggregation
  • The three aggregation operators: hash, object hash and sort
  • Why COUNT(DISTINCT) is more expensive than COUNT
  • How to spot a slow aggregation in a plan

Key takeaways

  • Aggregations run partial per task, shuffle, then final.
  • Partial aggregation shrinks data when keys repeat inside a task.
  • HashAggregate is fastest; ObjectHashAggregate handles object buffers; SortAggregate adds a sort.
  • Distinct aggregates need extra aggregation levels; several of them add an Expand.
  • approx_count_distinct is a mergeable sketch: one shuffle, bounded memory.

Check yourself

3 questions

What does partial_sum in a plan mean?

Show the answer

Each task's pre-aggregated sum, computed before the shuffle. Partial results are shuffled and merged by the final aggregation.

Which operator does collect_list typically use?

Show the answer

ObjectHashAggregate. Its buffer is a JVM object (a growing list), which HashAggregate cannot hold.

Why does COUNT(DISTINCT x) need a second shuffle?

Show the answer

Values must first be deduplicated across tasks by (key, x), then counted per key. A distinct count cannot be pre-aggregated into one number per task.

Keep going

Up next · lesson 23 of 30 · 4 min read
Serialization: Java, Kryo, Tungsten and Arrow
Where Spark serializes data and code, and which serializer matters for RDDs, DataFrames and Python.

Related lessons

Previous: Shuffle internals

Primary sources: Performance tuning · EXPLAIN