tf.data pagination with .skip()/.take() – unresolved shuffle ordering behavior
0 reputation · 07 Dec 2025, 03:22 UTC
0 reputation · 07 Dec 2025, 03:22 UTC
Goal: Retrieve a fixed window from a large dataset for training or inference, e.g., page 3 of size 100 via .skip(200).take(100). This approach is documented in the tf.data API.
When a shuffle operation precedes the pagination steps, the resulting page contains a random sample of the original data. If shuffle follows paging, the page is a contiguous slice of an already shuffled stream. The tf.data documentation does not specify which ordering preserves epoch‑level determinism or how the shuffle buffer is managed across epochs when combined with .skip() and .take(). Additionally, applying .take() to an unbounded source can result in an unknown cardinality, potentially breaking downstream operations that expect a fixed size.
In tf.data, the order of operations determines whether pagination is stable across epochs. Because .shuffle() creates a randomized buffer, its placement relative to .skip() and .take() fundamentally changes the resulting dataset.
.shuffle().skip(n).take(m), the pagination is non‑deterministic by default. Every time the dataset is re‑initialized (e.g., at the start of a new epoch), the shuffle buffer is re‑filled and re‑randomized. Consequently, .skip(200) will discard a different set of 200 elements each time, and your "page" will contain different data..skip(n).take(m).shuffle(), you are taking a deterministic slice of the original source and then randomizing only that slice. This ensures the page contents are the same every epoch, but the order within that page varies.The buffer_size in .shuffle() does not guarantee uniqueness across pages; it only controls the degree of randomness. If the buffer size is smaller than the total dataset, the shuffle is partial. To ensure a paged window contains unique, non‑repeating elements across the entire dataset without overlap between pages, you must avoid shuffling during the pagination process or use a fixed seed.
When .take() is applied to an unbounded source, the dataset's cardinality becomes tf.data.UNKNOWN_CARDINALITY. Downstream operations that require a known size (like certain batching configurations or progress bars) will fail to calculate the total steps.
To reliably determine cardinality after .take(m), you can manually track the count or, if the source is a known size, use the tf.data.Dataset.cardinality() method after the .take() operation, which should return m if the source had at least m elements.
If your goal is to have a "shuffled but stable" pagination (where Page 3 is always the same set of random elements across epochs), use .cache():
# Stable random pagination
# Shuffle once per dataset instance
dataset = dataset.shuffle(buffer_size=1000, seed=42)
dataset = dataset.cache()
# Now skip and take will be deterministic across epochs
page = dataset.skip(200).take(100)
Diagnostic Detail: Are you initializing a new tf.data.Dataset object for every page request, or are you maintaining a single iterator across the session?
Use comments to ask for clarification. Post a solution as an answer.
29,275 reputation · 07 Dec 2025, 06:33 UTC
To guarantee that the same page appears in every epoch, set a fixed seed and disable reshuffling:
ds = ds.shuffle(buffer_size=1000, seed=42, reshuffle_each_iteration=False)
.skip(200)
.take(100)
With reshuffle_each_iteration=False, TensorFlow uses the same permutation for each epoch, so .skip() and .take() always pick the same slice. If you omit the flag, the shuffle buffer is refilled each epoch and the page varies, even with a seed.
Note: a buffer size equal to the full dataset yields a true random permutation but can be memory‑intensive for large datasets.