skip to content

In tf.data, what does shuffle(buffer_size) actually do, and why does its position matter?

level: middleimportance: must knowfreq 78%

answer

  1. a sliding window, not a global permutation
  2. buffer_size elements held in memory
  3. sorted input plus small buffer equals single-class batches
  4. order versus batch, order versus repeat
  5. shuffle the shard list too

basics

~20 s

tf.data's shuffle keeps a buffer of buffer_size elements, emits one at random from it, and refills the slot from the source. It is a sliding-window shuffle, not a global one, so a small buffer over sorted data barely mixes anything.

solid answer

~50 s

`Dataset.shuffle(buffer_size)` fills an internal buffer with the first `buffer_size` elements, then repeatedly picks one uniformly at random, emits it, and pulls the next source element into the freed slot. So randomness is bounded by the window: with data sorted by label and a buffer of 1,000 over a million rows, the first batches can only contain elements from the first ~1,000 rows, which are all the same class. A truly uniform shuffle needs a buffer as large as the dataset, which is usually impossible, so in practice you shuffle the file list first and use a moderate element buffer on top. Position matters twice. Shuffle before `batch`, or you only permute the order of fixed batches while their contents never change. And prefer `shuffle().repeat()` over `repeat().shuffle()`: the former gives clean epoch boundaries where every element appears once before any repeats, while the latter lets the next epoch's elements bleed into the current one.

code

python · 16 lines
python
import tensorflow as tf

files = tf.data.Dataset.list_files("/data/train-*.tfrecord")  # shuffles files by default
ds = files.interleave(
    tf.data.TFRecordDataset,
    cycle_length=8,
    num_parallel_calls=tf.data.AUTOTUNE,
)
ds = (
    ds.shuffle(10_000)          # element-level shuffle, before batch
      .batch(32)
      .prefetch(tf.data.AUTOTUNE)
)

# fixed order for a validation split
val = tf.data.Dataset.range(100).shuffle(100, seed=0, reshuffle_each_iteration=False)

go deeper

for a junior

Know that shuffle needs a buffer_size argument, that bigger means better mixing, and that it belongs before batch so batch contents actually change between epochs.

for a middle

Explain the fill-emit-refill buffer algorithm, the RAM cost of the buffer, and why sorted input with a small buffer yields nearly single-class batches.

for a senior

Show the production pattern: shuffle shards and interleave them for long-range mixing, keep a modest element buffer, and know that cache after shuffle freezes the order and that a huge buffer adds first-step latency.

for a principal

Own the tradeoff between paying for randomness at write time and at read time — records written in random order once make every future job cheaper, and that decision belongs with the dataset format, not with each training script.

## The algorithm `tf.data.Dataset.shuffle(buffer_size, seed=None, reshuffle_each_iteration=None)` implements a reservoir-style sliding shuffle: 1. Pull `buffer_size` elements from the upstream dataset into a buffer. 2. Pick an index in the buffer uniformly at random, emit that element. 3. Pull the next upstream element into the vacated slot. 4. Repeat; at the end of the source, drain the buffer randomly. Two consequences fall straight out of this. First, memory cost is `buffer_size` elements held live — for decoded 224x224x3 float32 images that is roughly 600 KB each, so a buffer of 10,000 is about 6 GB of host RAM. Second, the shuffle is *local*: an element can only move forward or backward by roughly the buffer size in the output order. `buffer_size=1` is a no-op. ## Why locality bites Realistic datasets are rarely stored in random order. Records arrive sorted by class, by capture date, by user id, or grouped one class per file. Feeding a nearly-sorted stream into a model produces batches with almost no class diversity, and with batch normalization or any per-batch statistic that is actively harmful: the loss oscillates, gradients point in wildly different directions from batch to batch, and the run looks like a bad learning rate rather than a bad shuffle. The standard mitigation is two-level shuffling. Shuffle the *file list* — `Dataset.list_files(pattern)` shuffles by default, and `TFRecordDataset` shards can be interleaved in random order — then apply an element-level `shuffle` on top. Randomizing which shards you read from gives long-range mixing cheaply, and the element buffer handles local mixing. If you also control how the data is written, writing records in random order once, offline, makes every epoch cheaper forever. ## Ordering: shuffle before batch `ds.shuffle(1000).batch(32)` shuffles examples and then groups them, so batch composition is random and changes every epoch. `ds.batch(32).shuffle(1000)` groups first, so the batches are fixed sets of neighbouring examples and shuffling only permutes the order in which those fixed groups arrive. That is strictly weaker regularization, and if the data was sorted the batches stay single-class no matter how large the shuffle buffer is. The one time batch-then-shuffle is deliberate is when elements are already random and you want to cheaply reorder large pre-built batches. ## Ordering: shuffle before repeat `ds.shuffle(n).repeat()` completes one full pass — every element exactly once — before starting the next, giving clean, countable epochs. `ds.repeat().shuffle(n)` shuffles across the epoch boundary: the buffer can hold elements from pass 1 and pass 2 at the same time, so an element may be seen twice before another is seen once. It smooths the transition and can be marginally faster, but it makes 'one epoch' a fiction, which matters when you evaluate or checkpoint per epoch. If you use `steps_per_epoch` with an infinite dataset, be explicit about which you chose. ## reshuffle_each_iteration and reproducibility `reshuffle_each_iteration` defaults to `True`: each time you iterate the dataset again, you get a different permutation, which is what you want across epochs. Setting it to `False` together with a fixed `seed` gives the same order every pass — useful for a validation set or for a bit-reproducible run. Note that reproducibility also depends on the rest of the pipeline: `num_parallel_calls` with `deterministic` left at its default preserves order, but explicitly setting `deterministic=False`, or setting `tf.data.Options().deterministic = False`, trades order stability for throughput. ## Interactions to watch - **cache before shuffle, not after.** `ds.shuffle(n).cache()` caches one particular shuffled order and replays it forever, silently killing the reshuffle. Put `cache()` upstream of `shuffle`. - **Sharding for multi-worker input** should be applied deliberately relative to shuffle; sharding a shuffled stream by element gives each worker a random subset, sharding files gives each worker whole shards. - **Cost.** The buffer fill happens before the first element is emitted, so a huge buffer adds visible startup latency to every epoch and shows up as a slow first step. ## What a good answer sounds like Describe the buffer mechanically, state the memory cost, admit that it is not a global shuffle, and name the two ordering rules — before `batch`, before `repeat` — plus the file-level shuffle that makes a modest buffer sufficient.

  • How large a shuffle buffer would you actually pick for a 2-million-image dataset?
    Not two million — that would need hundreds of gigabytes of RAM. Pick a buffer that fits comfortably in host memory, often a few thousand decoded images, and get long-range mixing from randomizing the shard order plus interleaving many shards at once. If batches still look correlated, write the records in random order offline once instead of paying for it every epoch.
  • Why is ds.shuffle(n).cache() usually a bug?
    cache() records whatever it sees on the first full pass and replays exactly that on later passes. Placed after shuffle, it freezes one permutation, so reshuffle_each_iteration never takes effect and every epoch sees identical batches. Cache the expensive deterministic work upstream and let shuffle run after the cache.
  • You need bit-reproducible input order for a debugging run. What do you set?
    Pass an explicit seed to shuffle and set reshuffle_each_iteration=False, then make sure nothing downstream reorders: leave deterministic at its default rather than setting deterministic=False on map or interleave, and avoid tf.data.Options with deterministic disabled. Parallelism is still fine — determinism costs some throughput, not parallelism.
  • What is the difference in epoch semantics between shuffle().repeat() and repeat().shuffle()?
    shuffle().repeat() finishes a complete pass before repeating, so every element appears exactly once per epoch. repeat().shuffle() lets the buffer mix elements from consecutive passes, so an element can be emitted twice before another appears at all. The second smooths epoch boundaries but makes per-epoch evaluation and counting approximate.

saying these in an interview costs you the question

  • Claiming shuffle produces a uniform permutation of the whole dataset
  • Calling batch before shuffle and expecting mixed batches
  • Ignoring that the buffer holds buffer_size decoded elements in RAM
  • Placing cache after shuffle and wondering why epochs look identical
  • Thinking buffer_size=1 still shuffles a little

context