skip to content

Where should cache() go in a tf.data pipeline that does random augmentation?

level: seniorimportance: should knowfreq 50%

answer

  1. a line through the pipeline
  2. above the line runs once
  3. below the line runs every epoch
  4. frozen augmentation is silent
  5. in-memory by default, filename for disk

basics

~20 s

Put cache() after expensive deterministic work such as decode and resize, and before shuffle and random augmentation. Caching after augmentation freezes one augmented copy of each example and replays it every epoch, silently removing the augmentation's value.

solid answer

~50 s

`Dataset.cache()` records the elements it sees during the first complete pass and replays them from the cache on every later pass, skipping everything upstream. That makes placement a correctness question, not just a speed one. Anything upstream of `cache()` runs once, ever; anything downstream runs every epoch. So the cache belongs after the deterministic heavy lifting — file read, JPEG decode, resize, normalize — and before `shuffle` and any random augmentation. Cache after `shuffle` and you freeze a single permutation, defeating `reshuffle_each_iteration`; cache after augmentation and every epoch sees the identical augmented image, so you are training on a fixed enlarged dataset rather than a stream of fresh variations. Also decide where the cache lives: `cache()` with no argument holds every element in host memory, which OOMs on a large dataset, while `cache(filename)` writes a cache file to disk and trades RAM for I/O.

code

python · 9 lines
python
import tensorflow as tf

ds = (files_ds
      .map(decode_and_resize, num_parallel_calls=tf.data.AUTOTUNE)  # deterministic
      .cache()                       # or cache("/scratch/train.cache") for disk
      .shuffle(10_000)               # must stay below the cache
      .map(augment, num_parallel_calls=tf.data.AUTOTUNE)  # random, per epoch
      .batch(32)
      .prefetch(tf.data.AUTOTUNE))

go deeper

for a junior

Know that cache() makes later epochs faster by reusing the work already done, and that it must be placed before any random augmentation so each epoch still varies.

for a middle

Explain the line-through-the-pipeline model — upstream runs once, downstream runs every epoch — and the two silent bugs it causes when placed after shuffle or after augmentation.

for a senior

Size the cache before using it: estimate post-decode bytes per element, choose in-memory versus cache(filename), keep uint8 through the cache, and recognize that an incomplete first pass leaves no usable cache.

for a principal

Decide where preprocessing belongs across the whole fleet — deterministic work repeated by every job on every host is a stored-format problem, and pre-resized shards beat a per-job cache once more than a couple of teams read the dataset.

## What cache() actually records `tf.data.Dataset.cache(filename='')` sits at a point in the pipeline and memorizes the stream of elements passing through it during the first full iteration. On subsequent iterations it serves those elements directly and the upstream transformations are not executed at all. The mental model that keeps you out of trouble: **cache draws a line through the pipeline. Above the line, work happens once. Below the line, work happens every epoch.** ## The correctness trap Augmentation exists to show the model a different version of each example every epoch. If augmentation sits above the line, it runs once, and its single output is what every epoch replays: ``` # WRONG ds = ds.map(decode).map(augment).cache().shuffle(1000).batch(32) ``` Nothing raises. Training runs. The model simply sees a fixed dataset of N pre-augmented images instead of an effectively larger stream, so the regularization you budgeted for is gone and validation loss diverges from what your ablation predicted. The same applies to `shuffle`: `ds.shuffle(n).cache()` records one permutation and replays it, so `reshuffle_each_iteration=True` becomes a no-op and every epoch sees identical batches in identical order. The order that works: ``` ds = (ds .map(decode_and_resize, num_parallel_calls=tf.data.AUTOTUNE) # deterministic, expensive .cache() # the line .shuffle(10_000) .map(augment, num_parallel_calls=tf.data.AUTOTUNE) # random, per epoch .batch(32) .prefetch(tf.data.AUTOTUNE)) ``` The expensive, invariant part is paid once; the parts that must vary stay below the line. ## Memory versus disk `cache()` with no filename is an in-memory cache. Every element is retained in host RAM as it flows through, so the footprint is the entire post-decode dataset. For 100,000 images decoded to 224x224x3 float32 that is roughly 60 GB — you will OOM, often mid-first-epoch and confusingly far from the call site. Two ways out: - **`cache("/path/to/cachefile")`** streams the elements to a file on disk and reads them back on later epochs. This is often a large win when the upstream cost is decode/resize CPU work and the cached representation is smaller or cheaper to read than the source. Watch the disk footprint: caching *decoded* images can be far larger than the original JPEGs. - **Cache earlier, or cache less.** Cache after parsing but before decoding, or store `uint8` rather than `float32` and cast below the line. Halving the dtype halves the cache. ## The first pass must complete The cache is only usable once a full pass through it has finished. If you interrupt the first epoch, or place `take(k)` downstream and never consume the rest, the cache is incomplete and the upstream work re-runs. A common surprise is that the first epoch is *slower* than the no-cache baseline (it does the work and writes the cache) and epochs two onward are dramatically faster. For a one-epoch job, `cache()` is pure overhead. ## When cache is the wrong tool - **Datasets larger than RAM and larger than the local disk budget.** Then the answer is a cheaper source format, not a cache. - **Streaming or ever-changing data.** A cache pins the data as of the first pass; new records will never appear. - **Work that is already cheap.** Caching a pipeline whose upstream is a memory-resident tensor buys nothing and costs memory. - **Work that should be offline.** If the same deterministic decode/resize runs across many jobs and many hosts, do it once when writing the records and ship pre-resized shards; a per-job cache is a per-job tax. ## What to say in the interview Name the line-through-the-pipeline model, give the canonical order (read → deterministic map → cache → shuffle → random map → batch → prefetch), state the two silent bugs (frozen augmentation, frozen shuffle order), and mention the in-memory versus on-disk choice with a concrete size estimate. Then add the judgment: repeated deterministic preprocessing across many runs is a signal to fix the stored format, not to add a cache to each script.

  • Your cached pipeline OOMs the host during the first epoch. What do you change?
    Either move the cache to disk with cache(filename) or shrink what is cached. Caching decoded float32 images is the usual culprit — keep uint8 through the cache and cast below it, or cache after parsing but before decode. If neither fits, drop the cache and instead write pre-resized shards offline so the expensive step never runs in the training job.
  • Why is the first epoch slower with cache() than without it?
    The first pass does all the upstream work and additionally writes every element into the cache, in memory or to a file. Only from the second pass does the pipeline skip the upstream stages. That is why cache is worthless for a single-epoch job and why you should not benchmark it on epoch one.
  • What happens if training is interrupted partway through the first epoch?
    The cache never finalizes, so the next run repeats the upstream work rather than reading a partial cache. The same applies if a downstream take() or an early break means the pass never reaches the end of the dataset — the cache is only usable after a complete iteration through it.
  • When would you deliberately cache after shuffle?
    Essentially never for training, because it pins one permutation and kills reshuffling. The defensible case is a fixed evaluation set where you want a stable, reproducible order across runs — and even there, an explicit seed with reshuffle_each_iteration=False expresses the intent more clearly than relying on cache placement.

saying these in an interview costs you the question

  • Placing cache after random augmentation and expecting varied epochs
  • Calling cache() on a dataset far larger than host RAM
  • Assuming cache() writes to disk by default
  • Caching after shuffle and still expecting reshuffling each epoch
  • Judging cache's benefit from first-epoch timings

context