ML Engineer MasterClass (October) | 4 seats left

Spark · Partitions and tasks
AmazonAmazon Analytics
Amazon · How Spark Executes

Partitions and tasks

Inspect physical partitions and distinguish them from rows, task slots, and complete jobs.

Step 1 of 6 · Learn

Request physical partitions

repartition(2) redistributes this DataFrame into two physical partitions. getNumPartitions inspects the underlying RDD in this classic Spark runtime. Partition count is not row count, worker count, or a promise that all tasks run simultaneously. The tiny result DataFrame only displays that measured count; its own partition layout is unrelated.

Lesson reference: PySpark and Scala

Request physical partitions

repartition(2) redistributes this DataFrame into two physical partitions. getNumPartitions inspects the underlying RDD in this classic Spark runtime. Partition count is not row count, worker count, or a promise that all tasks run simultaneously. The tiny result DataFrame only displays that measured count; its own partition layout is unrelated.

PySpark example

from pyspark.sql import Window
from pyspark.sql.functions import lit

partitioned = orders.repartition(2)
result = spark.range(1).select(
    lit(partitioned.rdd.getNumPartitions()).alias("partition_count")
)
result.show(truncate=False)

Scala example

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

val partitioned = orders.repartition(2)
val result = spark.range(1).select(
    lit(partitioned.rdd.getNumPartitions).alias("partition_count")
)
result.show(false)

Inspect work per partition

mapPartitions visits a partition iterator. This example emits one row count for each partition, including empty ones, then summarizes those counts. A stage that processes all those partitions normally schedules one task per partition; retries and additional stages can produce more task attempts. The summary is not a measurement of total job tasks. Materializing an iterator with list is acceptable only for these six teaching rows.

PySpark example

from pyspark.sql import Window
from pyspark.sql.functions import sum, count

partitioned = orders.repartition(2)
# Emit one count per partition, including empty partitions.
counts = partitioned.rdd.mapPartitions(
    lambda partition: [(len(list(partition)),)]
)
partition_counts = spark.createDataFrame(counts, "rows_in_partition LONG")
result = partition_counts.agg(
    count("*").alias("partition_count"),
    sum("rows_in_partition").alias("row_count")
)
result.show(truncate=False)

Scala example

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

val partitioned = orders.repartition(2)
// Emit one count per partition, including empty partitions.
val partition_counts = partitioned.rdd.mapPartitions(
    partition => Iterator(partition.size.toLong)
).toDF("rows_in_partition")
val result = partition_counts.agg(
    count("*").alias("partition_count"),
    sum("rows_in_partition").alias("row_count")
)
result.show(false)

Keep empty partitions visible

A filter removes rows but does not itself redistribute surviving rows. After the explicit repartition, some partitions may become empty. Counting one summary per partition still includes them; counting distinct spark_partition_id values on surviving rows would not. This distinction matters when reasoning about work and skew.

PySpark example

from pyspark.sql import Window
from pyspark.sql.functions import col, sum, count

partitioned = orders.repartition(3).filter(col("status") == "Shipped")
# Emit one count per partition, including empty partitions.
counts = partitioned.rdd.mapPartitions(
    lambda partition: [(len(list(partition)),)]
)
partition_counts = spark.createDataFrame(counts, "rows_in_partition LONG")
result = partition_counts.agg(
    count("*").alias("partition_count"),
    sum("rows_in_partition").alias("row_count")
)
result.show(truncate=False)

Scala example

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

val partitioned = orders.repartition(3).filter(col("status") === "Shipped")
// Emit one count per partition, including empty partitions.
val partition_counts = partitioned.rdd.mapPartitions(
    partition => Iterator(partition.size.toLong)
).toDF("rows_in_partition")
val result = partition_counts.agg(
    count("*").alias("partition_count"),
    sum("rows_in_partition").alias("row_count")
)
result.show(false)
example.pyPySpark
1. Input rows6 orders; physical row placement is not fixed.
2. Repartition exchangeRedistribute rows into the requested number of partitions.
3. Downstream stageOne task per partition for a stage that processes all partitions; limited slots run tasks in waves.
Source orders6 rows
order_idcustomer_idstatustotalitem_count
1001101Delivered89.52
1002102Shipped1493
1003101Cancelled351
1004103Delivered219.994
1005104Delivered49.991
1006105Shipped1202
Request physical partitions
Result1 rows
partition_count
2
The example measures two partitions of orders. It does not infer their row placement or count tasks for the whole job.

The example is loaded in the editor. Run it as written, then try a small change.

Runs on the Spark backend. First startup may take a moment.

Run your code to see the result.