A Dataset in Spark is a distributed collection with a schema and a compile-time type. In Scala and Java you can write Dataset[Order] and transform it with ordinary functions over Order objects, and the compiler checks field names and types for you. The untyped DataFrame is simply Dataset[Row]. Python and R have no typed Dataset at all; everything there is a DataFrame.

The relationship between the two views, how encoders convert objects to Spark's internal binary rows, and why a typed lambda is often slower than a column expression are covered in Spark DataFrame and Dataset. This article picks up where that one stops: the typed operators themselves. It covers what groupByKey, mapGroups, reduceGroups, joinWith and the Aggregator API do, what physical plan each produces, which ones scale and which ones quietly move every row across the network, and then works through a sessionisation job in both styles.

Advertisement

The mental model: two worlds and a border

Spark executes on rows stored in its internal binary format, which Catalyst can reason about and whole-stage code generation can compile into tight loops. A typed operation needs real JVM objects, so Spark inserts a border crossing: DeserializeToObject turns rows into objects, your function runs, and SerializeFromObject turns the results back into rows. You can see these nodes in explain() output whenever you use a lambda.

Everything in this article follows from that border. Catalyst cannot look inside your lambda, so it cannot push a typed filter down into a Parquet scan, cannot prune columns your lambda does not read, and cannot reorder it with other operators. Column expressions have none of those limits. The working rule is therefore not "typed is bad" but "cross the border on purpose, as few times as possible, at the point where object code earns its keep".

case class Order(orderId: String, customerId: String, amount: Double, country: String)

val orders = spark.read.parquet("s3://shop/orders").as[Order]

// Typed filter: opaque to the optimizer, every column is read and deserialized.
orders.filter(o => o.country == "DE").explain()

// Column filter on the same Dataset: pushed into the Parquet scan.
orders.filter($"country" === "DE").explain()

Both lines return Dataset[Order]. Only the second shows PushedFilters: [IsNotNull(country), EqualTo(country,DE)] in the scan node. Typed and untyped operations mix freely on the same Dataset, which is the point: you keep the static type while using expressions for the parts the optimizer understands. Reading plans like these is covered in Spark explain plans.

Row-at-a-time operators: map, flatMap, mapPartitions

map and flatMap apply a function per element and need an encoder for the output type, which Scala supplies implicitly for case classes and primitives through spark.implicits._. Java callers pass one explicitly, such as Encoders.bean(Out.class) or Encoders.STRING().

mapPartitions hands you an iterator over a whole partition. Use it when each element would otherwise pay a fixed setup cost: opening a client, loading a model, compiling a regular expression. The iterator must be consumed lazily, because materialising it into a list holds the whole partition in memory.

val enriched: Dataset[Enriched] = orders.mapPartitions { it =>
  val client = GeoClient.connect(endpoint)            // once per partition, not per row
  val out = it.map(o => Enriched(o, client.lookup(o.country)))
  // Close only after the iterator is exhausted, or rows after the first batch fail.
  closeWhenDone(out)(client.close())                  // small helper of your own, see below
}

def closeWhenDone[T](it: Iterator[T])(cleanup: => Unit): Iterator[T] = new Iterator[T] {
  private var closed = false
  def hasNext: Boolean = { val h = it.hasNext; if (!h && !closed) { closed = true; cleanup }; h }
  def next(): T = it.next()
}

Spark has an internal class that does the same job, but it is not a public API, so a few lines of your own are safer across versions. The wrapper only closes on the success path; also register a listener with TaskContext.get().addTaskCompletionListener so the client is closed if the task fails midway.

Advertisement

groupByKey and KeyValueGroupedDataset

groupByKey(f) computes a key for every element with your function and returns a KeyValueGroupedDataset[K, V]. Unlike the untyped groupBy, it is not an aggregation by itself; it is a handle offering several ways to process groups, and they behave very differently at scale.

OperationWhat you writePartial aggregation before shuffle?Memory risk
mapGroups(K, Iterator[V]) => UNo: every row is shuffledOne task iterates a whole group
flatMapGroups(K, Iterator[V]) => IterableOnce[U]NoSame; worse if you buffer the iterator
reduceGroups(V, V) => VYesLow
agg(typedColumn)Aggregator.toColumn, up to severalYesBuffer per key
mapValuesV => W, before another group operationNot applicableNone
count()Rows per keyYesLow
cogroupTwo grouped datasets by the same key typeNo: both sides shuffledBoth groups for a key in one task
flatMapGroupsWithStateStreaming, per-key stateNoState store size
groupByKey paths: which typed operations get map-side partial aggregationDataset[Click]rows in Tungsten formatgroupByKey(_.userId)AppendColumns: key per rowdeserializereduceGroups / agg(Aggregator)partial agg before shufflemapGroups / flatMapGroupsevery row shuffledExchange (shuffle)hash partition on keypartial buffersall rowsFinal aggregate orsort by key + iterate groupDataset[Out]serialized back to rowsserializeGreen path: shuffle carries one small buffer per key per partition.Red path: shuffle carries every row, and one task must iterate the whole group.
reduceGroups and Aggregator-based agg combine values on the map side, so the shuffle carries small buffers. mapGroups and flatMapGroups ship every row and give one task the whole group.

The column that matters is the third. reduceGroups and agg with an Aggregator are planned as aggregates, so Spark computes a partial result per key inside each input partition, shuffles those small buffers, and merges them. mapGroups cannot do that, because your function needs to see the whole group at once; Spark must shuffle every row and sort by key so that one task sees one key's rows contiguously. For a key with a hundred million rows that is one task processing a hundred million rows, and adaptive skew-join handling does not help, because it targets joins, not group iteration. See adaptive query execution for what it does cover.

Rule of thumb: if the result per group is something you could compute incrementally, write it as reduceGroups or an Aggregator. Reach for mapGroups only when you genuinely need ordered or whole-group logic, and when you know the largest group fits comfortably in one task.

The Aggregator API

An Aggregator[IN, BUF, OUT] is a typed, mergeable aggregation: a zero value, a way to fold one input into a buffer, a way to merge two buffers, and a final step. Spark can then apply it partially on each partition and merge after the shuffle. Since Spark 3.0 the same class is also the recommended way to write untyped user-defined aggregate functions, registered with functions.udaf; the older UserDefinedAggregateFunction is deprecated.

import org.apache.spark.sql.{Encoder, Encoders}
import org.apache.spark.sql.expressions.Aggregator
import org.apache.spark.sql.functions.udaf

case class Stats(n: Long, sum: Double, max: Double)

object OrderStats extends Aggregator[Order, Stats, Stats] {
  def zero: Stats = Stats(0L, 0.0, Double.MinValue)
  def reduce(b: Stats, o: Order): Stats = Stats(b.n + 1, b.sum + o.amount, math.max(b.max, o.amount))
  def merge(a: Stats, b: Stats): Stats = Stats(a.n + b.n, a.sum + b.sum, math.max(a.max, b.max))
  def finish(b: Stats): Stats = b
  def bufferEncoder: Encoder[Stats] = Encoders.product[Stats]
  def outputEncoder: Encoder[Stats] = Encoders.product[Stats]
}

// Typed use: one Stats per customer, partial aggregation before the shuffle.
val perCustomer: Dataset[(String, Stats)] =
  orders.groupByKey(_.customerId).agg(OrderStats.toColumn.name("stats"))

// Untyped use from SQL or DataFrames, on an Aggregator over a single column.
object MaxAmount extends Aggregator[Double, Double, Double] {
  def zero = Double.MinValue
  def reduce(b: Double, a: Double) = math.max(b, a)
  def merge(x: Double, y: Double) = math.max(x, y)
  def finish(b: Double) = b
  def bufferEncoder = Encoders.scalaDouble
  def outputEncoder = Encoders.scalaDouble
}
orders.createOrReplaceTempView("orders")
spark.udf.register("max_amount", udaf(MaxAmount))
spark.sql("SELECT country, max_amount(amount) FROM orders GROUP BY country")

Two rules keep an Aggregator correct. merge must be associative and commutative, because Spark merges buffers in no particular order, and zero must be an identity for merge. The buffer should be small and fixed-size; an Aggregator whose buffer is a growing list is mapGroups in disguise, with the extra cost of serialising that list at every step.

joinWith versus join

join returns a flat DataFrame whose columns are the union of both sides, so a typed pipeline has to re-apply as[...] to a new case class afterwards. joinWith keeps both sides whole and returns Dataset[(A, B)]: each output element is a pair of the original objects.

case class Customer(customerId: String, tier: String)

val pairs: Dataset[(Order, Customer)] =
  orders.joinWith(customers, orders("customerId") === customers("customerId"), "inner")

val premium = pairs.filter($"_2.tier" === "gold").map { case (o, c) => o.amount }

The join condition is still a column expression, so join planning, broadcast selection and skew handling all work as usual. The costs are that each side becomes a nested struct column, which some downstream column references and writers handle less conveniently, and that for left outer joins the missing side arrives as null, not None, so pattern matching code must check for it explicitly.

Worked example: sessionising clickstream

Task: split each user's clicks into sessions, where a gap of more than 30 minutes starts a new session, and output one row per session with its start, end and click count. The input has a few hundred million clicks a day. Most users have tens of clicks; a handful of automated accounts have millions.

The typed version is natural to write:

case class Click(userId: String, ts: Long, url: String)
case class Session(userId: String, start: Long, end: Long, clicks: Int)

val gapMs = 30 * 60 * 1000L
val sessions = clicks.groupByKey(_.userId).flatMapGroups { (user, it) =>
  val sorted = it.toArray.sortBy(_.ts)          // buffers the whole group in memory
  val out = scala.collection.mutable.ArrayBuffer.empty[Session]
  var start = sorted.head.ts; var last = start; var n = 0
  for (c <- sorted) {
    if (c.ts - last > gapMs) { out += Session(user, start, last, n); start = c.ts; n = 0 }
    last = c.ts; n += 1
  }
  out += Session(user, start, last, n)
  out
}

It is correct and readable, and it fails on the automated accounts: the task holding a user with millions of clicks buffers them all in one array and either runs for an hour or dies with an out-of-memory error, while every other task finished in seconds. The column-expression version uses window functions, which Spark evaluates over sorted data in a spillable buffer:

import org.apache.spark.sql.expressions.Window
import org.apache.spark.sql.functions._

val w = Window.partitionBy($"userId").orderBy($"ts")
val sessions = clicks
  .withColumn("newSession", when(($"ts" - lag($"ts", 1).over(w)) > gapMs, 1).otherwise(0))
  .withColumn("sessionNo", sum($"newSession").over(w))
  .groupBy($"userId", $"sessionNo")
  .agg(min($"ts").as("start"), max($"ts").as("end"), count("*").cast("int").as("clicks"))
  .as[Session]                                   // back to the typed world at the end

The heavy user is still processed by one task, because a window partitioned by user cannot be split, but the window operator buffers rows in a structure that spills to disk past a threshold, so it slows down instead of failing. The final as[Session] restores the static type for the code that consumes the result, which is the pattern to copy: expressions in the middle, types at the edges. If you must keep the typed version, filter out or separately process keys above a size threshold first.

Streaming: flatMapGroupsWithState

In Structured Streaming, the same typed grouping supports arbitrary per-key state through mapGroupsWithState and flatMapGroupsWithState. Your function receives the key, the new values for that key in this micro-batch, and a GroupState[S] handle to read, update, remove or time out the state. Sessionisation is the canonical use, because the state per user is just the current session. Spark 4.0 added transformWithState as a more flexible successor with multiple state variables and timers; check the documentation for your version before relying on it. State storage, checkpointing and the state store providers are covered in Structured Streaming state.

Failure modes

  • Skewed groups in mapGroups or flatMapGroups. One task runs forever or runs out of memory. Use reduceGroups or an Aggregator, or a window function, and isolate giant keys.
  • Buffering the iterator. Calling toList or toArray on a group iterator turns a streaming pass into a memory bomb. Iterate once, or sort with expressions before grouping.
  • Typed filters before an expensive scan. A lambda filter disables predicate pushdown and column pruning, so the job reads every column of every file. Filter with columns first.
  • Encoder errors at runtime. Unsupported field types, such as some Java collection types or classes without a bean shape, fail when the encoder is derived. Keep Dataset types to case classes of primitives, strings, options, sequences, maps and nested case classes.
  • Schema drift breaking as[T]. as[T] resolves columns by name; a renamed upstream column fails the job, while an extra column is silently ignored until you write the Dataset out.
  • Non-associative Aggregator merge. Results differ between runs and cluster sizes. Unit test merge in both orders with random splits.

Trade-offs

The Dataset API buys compile-time checking of field access, the ability to express domain logic as plain methods, and code that unit tests without a SparkSession. It costs serialisation at each border crossing, optimizer blindness inside lambdas, and Scala or Java only. The memory and CPU costs of the border are explained with the binary row format in Tungsten. Teams that get the best of both write ingestion and heavy relational work as expressions, cross into typed code for the business rules that genuinely need it, and use Aggregators rather than group iteration whenever the logic allows.

What to do next

  1. Run explain() on your three most expensive typed jobs and count the DeserializeToObject and SerializeFromObject nodes.
  2. Replace typed filters and projections that precede a file scan with column expressions, and confirm PushedFilters appears in the plan.
  3. Search your code for mapGroups and flatMapGroups; for each, decide whether reduceGroups, an Aggregator or a window function can replace it.
  4. Measure the largest group per grouping key in production data before trusting any whole-group function.
  5. Port any UserDefinedAggregateFunction to an Aggregator registered with functions.udaf, and unit test merge for associativity.
  6. Adopt the edges pattern: read and transform with expressions, convert with as[T] where domain code begins, and convert back for writes.
Key takeaway: The typed Dataset API lets Scala and Java code operate on real objects, but every lambda is a border the optimizer cannot see through. Use column expressions for filtering, projection and relational work, cross into typed code only where domain logic needs it, prefer reduceGroups and Aggregators, which get partial aggregation, over mapGroups and flatMapGroups, which shuffle every row and hand one task a whole group, and check the physical plan to prove which path each job took.