Your Keras fit() on a repeated tf.data.Dataset never ends an epoch and ignores shuffle=True — why?
answer
- the pipeline owns batching and order
- who decides when an epoch ends?
- an infinite source has no end
- one argument raises, another is ignored
- you cannot slice an iterator
basics
~20 sA tf.data.Dataset already owns batching, shuffling and length, so fit() defers to it: the shuffle argument is ignored, batch_size is rejected, and an epoch runs until the dataset is exhausted — which after .repeat() never happens unless you pass steps_per_epoch.
solid answer
~40 sWhen `x` is a `tf.data.Dataset`, Keras treats it as an opaque producer of ready-made batches rather than as an array it can slice. Three consequences follow. `batch_size` must not be passed — the dataset already batches, and passing it raises. `shuffle=True` is ignored, so ordering comes only from `dataset.shuffle(buffer_size)` in the pipeline; a small or missing buffer means effectively unshuffled data with no warning. And an epoch ends when the dataset is exhausted, so a pipeline ending in `.repeat()` is infinite and the first epoch never completes — you must supply `steps_per_epoch` to define the epoch boundary yourself. `validation_split` is also unsupported for dataset inputs, because Keras cannot slice a dataset; pass a second dataset as `validation_data` (with `validation_steps` if it is also repeated).
code
python · 16 linesimport tensorflow as tf
import keras
x = tf.random.normal((256, 4))
y = tf.random.normal((256, 1))
ds = (
tf.data.Dataset.from_tensor_slices((x, y))
.shuffle(256) # shuffling belongs here, not in fit()
.batch(32) # batching belongs here too
.repeat() # infinite -> steps_per_epoch is mandatory
)
model = keras.Sequential([keras.layers.Dense(1)])
model.compile(optimizer="adam", loss="mse")
model.fit(ds, epochs=3, steps_per_epoch=8)go deeper
Remember that a tf.data pipeline, not fit(), decides batching and shuffling — call .shuffle(...) before .batch(...) in the pipeline and do not pass batch_size to fit().
Explain the contract: with a dataset input, shuffle is ignored, batch_size is rejected, validation_split is unsupported, and an epoch ends at dataset exhaustion. Derive from that why .repeat() requires steps_per_epoch.
Demonstrate the diagnosis, not just the rule. Show that you iterate one batch to check shapes and label order before blaming the model, and that you recognise a never-ending first epoch as a starved epoch boundary that also disables checkpointing and early stopping.
Own the guardrails: a house pipeline helper that shuffles before batching, asserts cardinality or requires an explicit step count, and validates the first batch at job start — so this class of silent data bug cannot reach a training run unnoticed.
## Why fit() behaves differently for a Dataset Keras `fit()` accepts several input species — NumPy arrays, a `tf.data.Dataset`, generators, `keras.utils.PyDataset` objects. With arrays, Keras owns everything: it can count samples, slice off a validation fraction, shuffle indices each epoch and cut batches of `batch_size`. With a `tf.data.Dataset` it owns none of that. A dataset is a pull-based iterator that yields whatever elements the pipeline was built to yield; Keras can only call `iter()` on it and consume batches. So the arguments that presuppose ownership are either rejected or ignored. This is the seam between the two halves of a TensorFlow stack, and it produces the classic pair of symptoms in the question. ## Symptom 1 — the epoch never ends For a dataset input with `steps_per_epoch=None`, an epoch runs **until the dataset is exhausted**. That is the correct behaviour for a finite pipeline. But a pipeline ending in `.repeat()` is infinite by construction, so exhaustion never arrives and the first epoch runs forever: the progress bar keeps counting steps, `epochs=10` never advances past epoch 1, and no callback that fires on epoch end ever runs — which also means checkpointing and early stopping never trigger. The fix is to define the boundary yourself with `steps_per_epoch`. The usual value is `ceil(num_samples / batch_size)`, computed from the raw data you know about, not from the dataset. The reason people add `.repeat()` in the first place is to avoid a partial final batch or an exhausted-iterator error under distribution; that is a legitimate pattern, and `steps_per_epoch` is its required companion. A related tell: because a repeated (or `from_generator`) dataset has unknown cardinality, Keras cannot show a total step count in the progress bar. A progress bar with no denominator is a hint that Keras does not know how long your epoch is. ## Symptom 2 — shuffle=True does nothing `fit(..., shuffle=True)` shuffles the sample order Keras itself generates. With a dataset, Keras generates nothing, so the argument is ignored for dataset (and generator) inputs. Shuffling has to happen in the pipeline: `dataset.shuffle(buffer_size)`. This is a silent-wrong-answer bug of the first order. Nothing raises. Training proceeds. The model just learns from data in file order, which for a dataset sorted by class produces long runs of a single label, unstable loss and a model that scores well in training and badly in production. The buffer size matters too: `shuffle` fills a buffer of that many elements and samples from it, so a buffer far smaller than the dataset gives only local shuffling — a common half-fix that looks like it works. Order in the pipeline also matters: shuffling after batching only permutes whole batches, leaving batch composition fixed. Shuffle before batch. ## The other arguments that change meaning - **`batch_size`** — must not be specified for a dataset input; Keras raises rather than silently ignoring it, because the dataset already produces batches. If your dataset is *not* batched, Keras treats each element as a whole batch, which usually surfaces as a shape mismatch in the first layer. - **`validation_split`** — unsupported for dataset inputs. Keras cannot slice off a fraction of an opaque iterator. Use `validation_data=val_ds`, having split the data upstream (different files, a filtered dataset, or a split before the pipeline is built). If the validation dataset is itself repeated, pass `validation_steps`. - **`class_weight`/`sample_weight`** — weighting still works, but for datasets the ergonomic route is to have the pipeline yield `(x, y, sample_weight)` tuples, which Keras understands directly. ## Diagnosing it in practice Before blaming the model, iterate the pipeline directly: take one batch, print its shapes and dtypes, and print the labels of the first few batches to see whether they are shuffled at all. Ninety percent of "Keras trains badly on my `tf.data` pipeline" issues are visible in that ten-line check — wrong batch shape, unshuffled labels, or an unbatched dataset. One more note for a Keras 3 answer: `tf.data.Dataset` is accepted as an input by Keras 3 on **all** backends, not only TensorFlow, and these same rules apply there. Using `tf.data` does not commit you to the TensorFlow backend for the model. ## How to answer Name the principle first — "the dataset owns batching, shuffling and length, so `fit()` defers to it" — then derive the three consequences. Interviewers ask this because it separates people who have actually shipped a `tf.data` training job from people who have only run `fit()` on NumPy arrays.
- How do you get a validation split when your input is a tf.data.Dataset?You split upstream and pass `validation_data` as a second dataset — separate files, a filtered or sharded dataset, or a split performed before the pipeline is built. `validation_split` is unsupported for dataset inputs because Keras cannot slice an opaque iterator. If the validation dataset is repeated, also pass `validation_steps` so evaluation terminates.
- What happens if you pass a dataset that was never batched?Keras treats each element of the dataset as one batch. An unbatched dataset therefore trains on batches of a single sample with a missing batch dimension, which usually raises a shape error in the first layer — or, worse, silently trains at batch size one and crawls. Call `.batch(n)` in the pipeline.
- Why can the progress bar show no total step count for a dataset input?Because the dataset's cardinality is unknown — typical after `.repeat()`, `filter()` or `from_generator`. Keras cannot compute a denominator, so it counts steps without a total. Treat a missing total as a signal to check whether you meant the epoch to be bounded by `steps_per_epoch`.
- Does any of this change when Keras 3 runs on the JAX or PyTorch backend?No. Keras 3 accepts `tf.data.Dataset` as an input on every backend and applies the same contract: the pipeline owns batching, shuffling and length. Using `tf.data` for input does not force the TensorFlow backend for the model itself.
saying these in an interview costs you the question
- Believes fit(shuffle=True) shuffles a tf.data.Dataset
- Passes batch_size alongside an already-batched dataset
- Uses validation_split with a Dataset input
- Adds .repeat() without setting steps_per_epoch
- Calls shuffle after batch and calls it shuffled