TensorFlow and Keras

tf.data input pipelines

tf.data turns "load, shuffle, transform, batch" into a conveyor belt that prepares the next batch while the model is still chewing the current one.

Read these first

On this page 5
  1. Why it exists
  2. How it works
  3. A real example you have seen
  4. Remember this
  5. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

tf.data is TensorFlow's conveyor belt for data — it loads, mixes, transforms and groups your examples, and keeps the next portion ready before the model asks for it.

Think of a busy kitchen with one cook and one helper. While the cook fries the current batch of pakoras, the helper is already chopping onions for the next. The cook never stands idle waiting for chopping. A tf.data pipeline is that helper: preparation overlaps with cooking.

Why it exists

A GPU or even a fast CPU model is a hungry cook. If data loading happens in plain Python — read a file, resize an image, wait — the expensive processor sits idle between batches. On real projects this is the most common reason training is slow, and the same disease exists in PyTorch.

tf.data fixes it by describing the whole preparation as a pipeline TensorFlow controls. Then TensorFlow can run steps in parallel, and prepare ahead.

How it works

Each step wraps the previous one, like sections of a conveyor belt:

source (your arrays or files)
   ↓ shuffle    — mix examples so no batch is all-one-kind
   ↓ map        — transform each example (resize, clean, encode)
   ↓ batch      — group into portions of, say, 32
   ↓ prefetch   — keep the NEXT portion ready during training
   → model

Two words to own here. Batch: the group of examples processed together in one training step. Prefetch: preparing portion N+1 while the model works on portion N — the helper-chops-ahead trick.

A real example you have seen

Video streaming apps buffer. While you watch minute 3, the app is already downloading minute 4, so playback never pauses. Remove buffering and every few seconds the video would stop to load. Prefetching in a data pipeline is buffering for a model — same idea, same reason.

Remember this

  • tf.data describes data preparation as a pipeline of wrapped steps.
  • Order matters: mix first, transform, then group, then prefetch.
  • The goal is overlap — the model should never wait for data.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install tensorflow

Outputs verified with TensorFlow 2.21, CPU. The shuffle order is seeded, so this exact output reproduces on one version; across versions the order may differ.

A complete pipeline on six examples

pipeline.py
import tensorflow as tf

# Six sentences with sentiment labels — stand-ins for a dataset of millions
texts = ["good", "great", "fine", "bad", "awful", "poor"]
labels = [1, 1, 1, 0, 0, 0]

ds = tf.data.Dataset.from_tensor_slices((texts, labels))

ds = (ds
      .shuffle(buffer_size=6, seed=0)      # shuffle BEFORE batching
      .map(lambda t, l: (tf.strings.length(t), l),
           num_parallel_calls=tf.data.AUTOTUNE)
      .batch(2)
      .prefetch(tf.data.AUTOTUNE))         # prepare the next batch during training

for lengths, labs in ds:
    print("lengths:", lengths.numpy(), "labels:", labs.numpy())
Output
lengths: [4 5] labels: [1 1]
lengths: [4 4] labels: [1 0]
lengths: [3 5] labels: [0 0]

The walkthrough

from_tensor_slices slices along the first dimension: six texts and six labels become six (text, label) pairs. Feeding model.fit(ds) from here works directly — Keras accepts a dataset in place of arrays.

shuffle(buffer_size=6) keeps a waiting-room of that many examples and picks randomly from it. With buffer size = dataset size, the mix is perfect. With a buffer of 2 on a million sorted rows, "random" picks come from a two-example window — barely shuffled at all. Undersized buffers are the silent killer of shuffling: symptom, batches full of one class.

map(...) transforms each pair — here to (text length, label), a stand-in for real preprocessing like tokenisation. num_parallel_calls=tf.data.AUTOTUNE lets TensorFlow choose how many examples to transform in parallel. The function inside map gets traced into a graph, so the tf.function rules apply — TensorFlow ops thrive; arbitrary Python needs the escape hatch tf.py_function, which costs the parallelism.

batch(2) stacks pairs into groups: notice the output tensors hold two values each. Three batches from six examples.

prefetch(AUTOTUNE) decouples producer from consumer. It has nothing to show in a print-out — its entire effect is timing. Last stage, always.

Common mistakes

Shuffling after batching. batch(2).shuffle(...) shuffles whole batches while each batch keeps its original neighbours. The model sees the same pairs together every epoch. Shuffle examples, then batch: order as shown above.

A tiny shuffle buffer on sorted data. Data sorted by class + shuffle(buffer_size=100) yields long single-class stretches. If your loss oscillates with a period, suspect this. Buffer as large as memory allows, or shuffle file order too.

Doing heavy work per epoch that could be cached. An expensive map reruns every epoch. Insert .cache() after it — before shuffle and batch — and epoch 2 onward reads from memory. For datasets bigger than RAM, .cache(filename) spills to disk.

Losing the dataset by forgetting reassignment. ds.batch(32) returns a new dataset; called bare, it does nothing. Every pipeline stage must be assigned: ds = ds.batch(32) or the chained style above.

Try it yourself

Move .shuffle(6, seed=0) to after .batch(2) and rerun several times without the seed. Watch the same pairs travel together in every run, in different orders — the broken version looks shuffled until you check the pairs.

What to learn next

Researcher — Mathematics and papers.

The pipeline as a streaming dataflow

A tf.data pipeline is a pull-based iterator graph: each stage implements GetNext() and requests elements from upstream on demand. Transformations declare their parallelism and the runtime schedules them over a shared threadpool. AUTOTUNE runs a gradient-free optimiser over the parallelism/buffer parameters, modelling the pipeline as a network of M/M/k queues and minimising expected output latency subject to a memory budget (Murray et al. 2021, tf.data: a machine learning data processing framework).

Throughput model

Steady-state step time with prefetching is $\max(t_{\text{prep}}, t_{\text{train}})$ rather than $t_{\text{prep}} + t_{\text{train}}$ — the entire benefit in one formula. Symbols: $t_{\text{prep}}$ — wall time to produce one batch; $t_{\text{train}}$ — model compute time per batch. Prefetch buffers of depth $b$ absorb variance: they smooth preparation-time jitter with burst tolerance growing in $b$, at memory cost $b \times$ batch bytes. Diagnosis rule: if accelerator utilisation rises when you lower model cost, the pipeline was never the bottleneck; if step time tracks $t_{\text{prep}}$, it is.

Shuffle-buffer statistics

Buffered shuffling with buffer size $k$ over $n$ examples is not a uniform random permutation for $k < n$. An element at source position $i$ can only appear in output positions $\ge i - k$: the permutation is $k$-local. Class-sorted inputs therefore retain long-range class autocorrelation for small $k$, biasing SGD — batch composition correlates with training step, violating the i.i.d. sampling assumption behind the mini-batch gradient estimate. Full uniformity requires $k = n$ or a pre-shuffled source (shuffle file shards, then records within a window).

Determinism trade-offs

Parallel map with deterministic=False (or interleaved file reads) returns elements in completion order, buying throughput at the price of run-to-run reproducibility — one of several places reproducibility quietly leaks away even with every seed set. The options().experimental_deterministic flag restores order at a latency cost.

References

  • Murray et al. (2021), tf.data: a machine learning data processing framework, VLDB.
  • Kuchnik et al. (2022), Plumber: diagnosing and removing performance bottlenecks in machine learning data pipelines.
  • The PyTorch counterpart design: DataLoader internals.

What to learn next