Menu

Apache Spark course · Lesson 2 of 5

Spark RDD Fundamentals: Transformations, Actions and Key-Value Operations

Learn Spark's RDD layer: lineage and partitions, transformations versus actions, map versus flatMap, reduceByKey versus groupByKey, mapPartitions and closures.

  • Beginner
  • 17 min read
  • Updated Oct 2026
On this page
  1. Sample data
  2. RDD fundamentals
  3. What it is
  4. How you create one
  5. You can read the lineage
  6. RDDs versus DataFrames today
  7. Pitfalls
  8. In interviews
  9. RDD transformations versus actions
  10. What it is
  11. Common operations
  12. Pitfalls
  13. In interviews
  14. map versus flatMap
  15. What it is
  16. Pitfalls
  17. In interviews
  18. reduceByKey versus groupByKey
  19. What it is
  20. The wider family
  21. Pitfalls
  22. In interviews
  23. mapPartitions
  24. What it is
  25. Pitfalls
  26. In interviews
  27. Closures and serialization
  28. What it is
  29. How to avoid trouble
  30. In interviews
  31. Practice questions
  32. Key takeaways

The RDD (Resilient Distributed Dataset) is the data structure the whole Spark engine is built on. You will write DataFrame code most of the time, but DataFrames compile down to RDDs, and every concept interviewers probe (lineage, laziness, narrow and wide dependencies, closures) is easiest to see at the RDD level. This lesson covers the RDD API you are expected to know and the mistakes it lets you make.

Sample data

A tiny text dataset split into two partitions, and a SparkSession in local mode. Later blocks build on these variables.

from pyspark.sql import SparkSession

spark = SparkSession.builder.master("local[2]").appName("rdd-fundamentals").getOrCreate()
sc = spark.sparkContext

lines = sc.parallelize([
    "spark makes big data simple",
    "big data needs big clusters",
    "",
    "spark spark spark",
], 2)

print(lines.getNumPartitions())
print(lines.glom().collect())   # glom() turns each partition into a list so you can see the split
2
[['spark makes big data simple', 'big data needs big clusters'], ['', 'spark spark spark']]

RDD fundamentals

What it is

An RDD is a read-only collection of records split into partitions that live on different executors. It has five defining properties in Spark’s source: a list of partitions, a function to compute each partition, a list of dependencies on parent RDDs, and optionally a partitioner (for key-value RDDs) and preferred locations (for data locality).

The key word is resilient. Spark does not replicate RDD data to survive failures. It records the lineage: the chain of transformations that produced the RDD. If an executor dies and a partition is lost, Spark re-runs just the steps needed to rebuild that one partition from its source.

Property What it means for you
Partitioned One task per partition; partition count sets parallelism
Immutable Transformations return a new RDD; the original never changes
Lazily evaluated Nothing runs until an action asks for a result
Lineage-based recovery Lost partitions are recomputed, not restored from copies
Typed by your code Records are any Python objects; Spark cannot see inside them

How you create one

  • sc.parallelize(collection, numSlices) distributes a driver-side collection (tests and demos).
  • sc.textFile(path) reads lines from files, one partition per file block or split.
  • df.rdd exposes a DataFrame as an RDD of Row objects.

You can read the lineage

toDebugString() prints the lineage. Indentation changes (+-) mark a shuffle dependency, which is where a new stage starts.

pairs = lines.flatMap(lambda line: line.split()).map(lambda w: (w, 1))
counts = pairs.reduceByKey(lambda a, b: a + b)
print(counts.toDebugString().decode())
(2) PythonRDD[6] at RDD at PythonRDD.scala:59 []
 |  MapPartitionsRDD[5] at mapPartitions at PythonRDD.scala:197 []
 |  ShuffledRDD[4] at partitionBy at DirectMethodHandleAccessor.java:103 []
 +-(2) PairwiseRDD[3] at reduceByKey at <block 2>:2 []
    |  PythonRDD[2] at reduceByKey at <block 2>:2 []
    |  ParallelCollectionRDD[0] at readRDDFromFile at PythonRDD.scala:326 []

RDD ids vary between runs, and <block 2>:2 is the source location of the reduceByKey call (your script name and line number when you run it yourself). The (2) is the partition count at each level. Everything below the +- runs in one stage, everything above it in the next.

RDDs versus DataFrames today

RDD records are opaque Python objects, so Spark cannot optimise them: no column pruning, no predicate pushdown, no code generation, and every record crosses between the JVM and Python. DataFrames carry a schema, so the Catalyst optimiser can do all of that (see Catalyst and Tungsten). Use DataFrames by default and drop to RDDs for genuinely unstructured processing or low-level control. Also note that Spark Connect clients have no RDD API at all, which is a growing reason to keep RDD code out of new pipelines.

Pitfalls

  • Using RDDs for tabular data because “it is lower level, so faster”. In PySpark it is usually slower.
  • Calling collect() on a large RDD: every record is pickled and sent to the driver.
  • Expecting an RDD to remember results: without cache() every action recomputes the lineage.

In interviews

“What does the R in RDD stand for, and how does Spark recover a lost partition?” A strong answer explains lineage-based recomputation, mentions that only the lost partitions are rebuilt, and adds that a long lineage can be cut with checkpointing.

RDD transformations versus actions

What it is

A transformation describes a new RDD from an existing one (map, filter, flatMap, reduceByKey, join). It returns immediately without touching data. An action asks for a result (collect, count, take, reduce, saveAsTextFile, foreach). Only actions make Spark run tasks.

def broken(x):
    return 1 / 0

risky = sc.parallelize(range(10)).map(broken)
print("defined, nothing failed yet")

try:
    risky.count()
except Exception as e:
    print("failed at the action:", type(e).__name__)
defined, nothing failed yet
failed at the action: PythonException

The bug is in the map, but the error only appears at count(). That is laziness in one example.

Common operations

Transformations (lazy) Actions (run a job)
map, flatMap, filter, mapPartitions collect, take(n), first
distinct, union, sample count, countByKey, countByValue
reduceByKey, groupByKey, aggregateByKey, sortByKey reduce, fold, aggregate
join, leftOuterJoin, cogroup saveAsTextFile, foreach, foreachPartition
repartition, coalesce, partitionBy takeOrdered, top
words = lines.flatMap(lambda line: line.split())
print(words.take(3), words.first())
print(words.distinct().count())
print(lines.filter(lambda l: "spark" in l).collect())
print(dict(pairs.countByKey()))
['spark', 'makes', 'big'] spark
7
['spark makes big data simple', 'spark spark spark']
{'spark': 4, 'makes': 1, 'big': 3, 'data': 2, 'simple': 1, 'needs': 1, 'clusters': 1}

Pitfalls

  • Each action re-runs the whole lineage. Two actions on an uncached RDD mean the source is read twice.
  • countByKey, countByValue and collectAsMap return a dictionary to the driver; use them only when the number of keys is small.
  • foreach runs on executors, so print inside it goes to executor logs, not your notebook.

In interviews

Be ready to classify operations quickly and to explain why Spark is lazy: it lets Spark see the whole chain before running, pipeline narrow steps into one task, and skip work whose results are never used. See also the transformation versus action question.

map versus flatMap

What it is

Both apply your function to every record. map produces exactly one output per input. flatMap lets each input produce zero, one or many outputs: your function returns an iterable, and Spark flattens it into the result.

as_lists = lines.map(lambda line: line.split(" "))
print(as_lists.collect())

flat = lines.flatMap(lambda line: line.split())
print(flat.collect())

print(as_lists.count(), flat.count())
[['spark', 'makes', 'big', 'data', 'simple'], ['big', 'data', 'needs', 'big', 'clusters'], [''], ['spark', 'spark', 'spark']]
['spark', 'makes', 'big', 'data', 'simple', 'big', 'data', 'needs', 'big', 'clusters', 'spark', 'spark', 'spark']
4 13

map kept four records (one list per line). flatMap produced thirteen words. Notice the empty line: split(" ") turned it into [''], a list with one empty string, while split() with no argument returned an empty list, so flatMap dropped that line entirely. That is how flatMap doubles as a filter: return an empty iterable to drop a record.

Pitfalls

  • Returning a string from a flatMap function flattens it into characters, because a string is iterable.
  • Returning None from flatMap raises an error; return [] instead.
  • The DataFrame equivalent of flatMap is explode on an array column.

In interviews

The classic test is word count: flatMap to words, map to (word, 1), reduceByKey to sum. Explain why map there would give you lists, not words.

reduceByKey versus groupByKey

What it is

Both work on key-value RDDs (records that are 2-tuples) and both cause a shuffle so that all values for a key meet in one partition. The difference is when values are combined.

  • reduceByKey(func) first combines values within each partition on the map side (a combiner), then shuffles one partial result per key per partition, then combines again.
  • groupByKey() shuffles every value and builds the full list of values for each key on the reduce side.
counts = pairs.reduceByKey(lambda a, b: a + b)
print(sorted(counts.collect(), key=lambda kv: (-kv[1], kv[0])))

grouped = pairs.groupByKey().mapValues(sum)
print(sorted(grouped.collect(), key=lambda kv: (-kv[1], kv[0])))
[('spark', 4), ('big', 3), ('data', 2), ('clusters', 1), ('makes', 1), ('needs', 1), ('simple', 1)]
[('spark', 4), ('big', 3), ('data', 2), ('clusters', 1), ('makes', 1), ('needs', 1), ('simple', 1)]

Same answer. The cost is very different. This block runs both on 200,000 records with 10 keys and reads the shuffle-write bytes from the Spark UI’s REST API:

import json, urllib.request

def shuffle_write_bytes(description):
    opener = urllib.request.build_opener(urllib.request.ProxyHandler({}))
    port = sc.uiWebUrl.rsplit(":", 1)[1]
    url = f"http://localhost:{port}/api/v1/applications/{sc.applicationId}/stages"
    return sum(s["shuffleWriteBytes"] for s in json.load(opener.open(url))
               if s.get("description") == description)

big = sc.parallelize(range(200_000), 4).map(lambda i: (i % 10, i))

sc.setJobDescription("reduceByKey")
big.reduceByKey(lambda a, b: a + b).collect()
sc.setJobDescription("groupByKey")
big.groupByKey().mapValues(sum).collect()
sc.setJobDescription(None)

print("reduceByKey shuffle write:", shuffle_write_bytes("reduceByKey"), "bytes")
print("groupByKey shuffle write: ", shuffle_write_bytes("groupByKey"), "bytes")
reduceByKey shuffle write: 1394 bytes
groupByKey shuffle write:  737328 bytes

reduceByKey shuffled one running total per key per partition; groupByKey shuffled all 200,000 values, roughly 500 times more data for the same result.

The wider family

Operation Map-side combine Use it for
reduceByKey(f) Yes Sums, counts, max/min: same type in and out, associative and commutative f
foldByKey(zero, f) Yes Like reduceByKey with a starting value
aggregateByKey(zero, seqOp, combOp) Yes Result type differs from the value type, such as (sum, count) for averages
combineByKey(create, merge, mergeCombiners) Yes The general form the others are built on
groupByKey() No When you truly need every value together (sorting a key’s events, building a full list)

An average per key with aggregateByKey, keeping a (sum, count) pair:

sales = sc.parallelize([("uk", 10.0), ("us", 20.0), ("uk", 30.0), ("in", 5.0), ("us", 40.0)], 2)
sum_count = sales.aggregateByKey(
    (0.0, 0),
    lambda acc, v: (acc[0] + v, acc[1] + 1),     # within a partition
    lambda a, b: (a[0] + b[0], a[1] + b[1]),      # across partitions
)
print(sorted(sum_count.mapValues(lambda t: t[0] / t[1]).collect()))
[('in', 5.0), ('uk', 20.0), ('us', 30.0)]

Pitfalls

  • groupByKey on a skewed key builds one huge list in one task, a common cause of executor out-of-memory errors.
  • The function passed to reduceByKey must be associative and commutative, because Spark applies it in any order across partitions. Subtraction or “keep the first value” gives non-deterministic results.
  • Averaging by reduceByKey on averages is wrong; carry sum and count, then divide.

In interviews

“Why prefer reduceByKey over groupByKey?” is near-certain. Say: map-side combining shrinks shuffle data; groupByKey sends every value across the network and holds all of them for a key in memory. Add the nuance that DataFrame groupBy().agg() already does partial aggregation automatically (the HashAggregate → Exchange → HashAggregate pattern in a plan), so the choice mainly matters in RDD code.

mapPartitions

What it is

map calls your function once per record. mapPartitions calls it once per partition, passing an iterator over that partition’s records, and expects an iterator back. Use it when each call has a fixed set-up cost: opening a database connection, loading a model, compiling a regular expression, creating an HTTP client.

The example counts how often an expensive lookup object is built:

setups = sc.accumulator(0)

class RateLookup:
    def __init__(self):
        setups.add(1)        # pretend this is a slow connection or model load
        self.rates = {"GBP": 1.0, "USD": 0.79, "EUR": 0.85}

    def to_gbp(self, amount, ccy):
        return round(amount * self.rates[ccy], 2)

payments = sc.parallelize(
    [(100.0, "USD"), (50.0, "EUR"), (20.0, "GBP"), (10.0, "USD"), (80.0, "EUR"), (5.0, "GBP")] * 2, 3)

per_row = payments.map(lambda p: RateLookup().to_gbp(*p)).collect()
print("set-ups with map:", setups.value)

setups.value = 0

def convert_partition(rows):
    lookup = RateLookup()            # once per partition
    for amount, ccy in rows:
        yield lookup.to_gbp(amount, ccy)

per_part = payments.mapPartitions(convert_partition).collect()
print("set-ups with mapPartitions:", setups.value)
print(per_part == per_row, per_part[:4])
set-ups with map: 12
set-ups with mapPartitions: 3
True [79.0, 42.5, 20.0, 7.9]

Twelve records cost twelve set-ups with map, but three partitions cost three with mapPartitions, with identical results. mapPartitionsWithIndex also passes the partition number, and foreachPartition is the action form for writing to external systems.

Pitfalls

  • Materialising the partition: list(rows) loads the whole partition into memory. Stream with a generator (yield) instead.
  • Forgetting to return an iterator: a function that returns None fails; one that returns a number fails because it is not iterable.
  • Leaking connections: close resources in a try/finally, since a task can fail mid-partition.
  • For DataFrames, the equivalent is mapInPandas or mapInArrow, which batch rows through Arrow.

In interviews

Expect “map versus mapPartitions” and “how would you write each row to an external database efficiently?”. The answer is foreachPartition with one connection per partition and batched writes, plus a note about idempotency because tasks can be retried.

Closures and serialization

What it is

When you pass a function to an RDD operation, Spark serializes the function together with every variable it references (its closure), ships the bytes to each executor, and deserializes a separate copy per task. PySpark uses cloudpickle for functions and pickle for data. Two consequences follow.

1. Changes made in a task never reach the driver. Each task mutates its own copy:

counter = 0

def increment(x):
    global counter
    counter += 1

sc.parallelize(range(100), 4).foreach(increment)
print("counter on driver:", counter)

total = sc.accumulator(0)
sc.parallelize(range(100), 4).foreach(lambda x: total.add(1))
print("accumulator:", total.value)
counter on driver: 0
accumulator: 100

In local mode, with some cluster setups, code like this can appear to work, which is why the RDD programming guide calls its behaviour undefined. Use an accumulator for counters, or better, compute the value with an action.

2. Everything in the closure must be serializable. Locks, open sockets, database connections, file handles and SparkContext itself cannot be pickled:

import threading

lock = threading.Lock()
try:
    sc.parallelize([1, 2]).map(lambda x: (x, lock)).collect()
except Exception as e:
    print(type(e).__name__, "-", str(e).splitlines()[0])
PicklingError - Could not serialize object: TypeError: cannot pickle '_thread.lock' object

How to avoid trouble

  • Create non-serializable resources inside the task (in mapPartitions), not on the driver.
  • Keep closures small. Referencing self.some_field inside a method ships the whole self object, including anything large or unpicklable on it. Copy the field to a local variable first.
  • For a large read-only lookup used by every task, use a broadcast variable so each executor receives it once rather than once per task.
  • Never reference spark or sc inside a transformation; executors cannot use them.
  • In Scala and Java the same issue appears as Task not serializable (java.io.NotSerializableException). spark.serializer defaults to Java serialization; Kryo (org.apache.spark.serializer.KryoSerializer) is faster and more compact for JVM objects in RDD code.

In interviews

Interviewers often show the counter example and ask what it prints. Explain that each task gets a deserialized copy of the closure, so the driver’s variable is untouched, then offer accumulators or an action-based alternative. Follow-ups: “What does Task not serializable mean?” and “Why would you broadcast a dictionary instead of capturing it?”

Practice questions

How does Spark rebuild a lost RDD partition, and when is that expensive?

Spark walks the partition’s lineage back to the nearest available data (the source, a cached copy, surviving shuffle files or a checkpoint) and re-runs only the transformations needed for that partition. It becomes expensive when the lineage is long or crosses shuffles whose files were also lost, because whole parent stages must rerun. Caching or checkpointing intermediate results shortens the recovery path.

You run rdd.map(parse).filter(valid) and nothing happens. Why, and what makes it run?

map and filter are transformations; they only build lineage. Spark runs tasks only when an action such as count(), take(), collect() or saveAsTextFile() requests a result. Any bug inside parse will surface at that action.

What is the output count of map versus flatMap over 4 lines where one line is empty and you split on whitespace?

map(lambda l: l.split()) always returns 4 records (one list per line, the empty line giving []). flatMap returns the total number of words, and the empty line contributes nothing because its list is empty.

Why is groupByKey followed by summing slower than reduceByKey, and when is groupByKey still right?

reduceByKey sums values inside each map partition before the shuffle, so only one partial sum per key per partition crosses the network. groupByKey ships every value and holds all values of a key in one task’s memory, risking spills or out-of-memory errors on hot keys. groupByKey is still right when you genuinely need all values together, for example to sort a user’s events or run a non-associative calculation, ideally after reducing the data first.

How would you compute the average value per key with RDDs?

Use aggregateByKey((0.0, 0), seqOp, combOp) (or combineByKey) to build a (sum, count) pair per key with map-side combining, then mapValues(lambda t: t[0] / t[1]). Averaging partial averages is wrong because partitions hold different numbers of records.

A job writes each record to a REST API with map and creates a client per record. How do you fix it?

Switch to foreachPartition (or mapPartitions if you need results) and create one client per partition, sending records in batches. Close the client in finally. Because tasks may be retried or speculatively duplicated, make the writes idempotent, for example by using upserts keyed on a record id.

What is a closure in Spark and what goes wrong with it?

The closure is the function plus the variables it references, serialized and sent to executors. Each task gets its own copy, so mutations do not return to the driver, and anything unpicklable (locks, connections, the SparkContext) makes the job fail at submission. Large captured objects are re-sent with every task; broadcast them instead.

Key takeaways

  • An RDD is partitioned, immutable and lazily evaluated, and it recovers lost partitions by recomputing them from lineage.
  • Transformations build lineage; actions run jobs, and each action recomputes uncached lineage.
  • map is one-to-one; flatMap is one-to-many and can drop records by returning an empty iterable.
  • Prefer reduceByKey or aggregateByKey to groupByKey: map-side combining cut shuffle bytes by about 500 times in the example.
  • Use mapPartitions or foreachPartition to pay set-up costs once per partition.
  • Closures are copied to every task: driver variables are not updated, and every captured object must be serializable.
  • Prefer DataFrames for structured data; RDDs bypass the optimiser and are not available through Spark Connect.

By Data Career Hub Editorial · Last reviewed Oct 2026 · All examples run on PySpark 4.2.0 in local mode (local[2]). Shuffle byte counts come from the Spark UI REST API of that run and will differ slightly on other machines.

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

Search
Filter by type