skip to content

How does tf.distribute shard a tf.data.Dataset across workers, and when does it fail?

level: seniorimportance: should knowfreq 45%

answer

  1. Two splits: across workers, then across GPUs
  2. AUTO tries files first
  3. Fewer files than workers is the trap
  4. The fallback reads everything and discards
  5. InputContext when you want to shard yourself

basics

~20 s

strategy.experimental_distribute_dataset auto-shards: with AUTO it splits by file when the source is file-based and there are at least as many files as workers, otherwise every worker reads all records and discards those not its own. Too few files kills throughput.

solid answer

~40 s

When you wrap a dataset with `strategy.experimental_distribute_dataset(...)`, `tf.data` auto-shards it across **workers** (the split across GPUs inside one worker happens separately, per batch). The policy lives at `tf.data.Options().experimental_distribute.auto_shard_policy` and defaults to `AUTO`, which tries `FILE` — give each worker a disjoint subset of the input files — and falls back to `DATA` if the pipeline is not file-based or has too few files. `DATA` sharding is correct but wasteful: every worker reads every record and keeps only `index % num_workers == id`, so your input I/O multiplies by the worker count. Setting the policy explicitly to `FILE` when there are fewer files than workers raises an error instead of degrading quietly. When you need control, use `strategy.distribute_datasets_from_function(fn)`: `fn` receives an `InputContext` with `num_input_pipelines`, `input_pipeline_id` and `get_per_replica_batch_size(...)`, and you do the sharding and batching yourself.

code

python · 14 lines
python
import tensorflow as tf

strategy = tf.distribute.MultiWorkerMirroredStrategy()

files = tf.data.Dataset.list_files("/data/train-*.tfrecord", shuffle=False)
dataset = tf.data.TFRecordDataset(files).batch(64)

options = tf.data.Options()
options.experimental_distribute.auto_shard_policy = (
    tf.data.experimental.AutoShardPolicy.FILE
)
dataset = dataset.with_options(options)

dist_dataset = strategy.experimental_distribute_dataset(dataset)

go deeper

for a junior

Know that a distributed dataset is split so different workers see different data, and that this is done for you when you wrap the dataset with the strategy.

for a middle

Name the policies — AUTO, FILE, DATA, OFF — say where they are set, and explain that FILE splits input files while DATA has every worker read everything and discard most of it.

for a senior

Diagnose the real cases: too few shards forcing the DATA fallback, uneven shard sizes stalling a synchronous job, and taking control with distribute_datasets_from_function and InputContext.

for a principal

Treat shard layout as a platform decision made at data-preparation time — target shard count and size, an explicit FILE policy so misconfiguration fails at start, and an input-throughput check in the training job's standard telemetry.

## Two different splits Distributed input in TensorFlow involves two separate divisions, and conflating them causes most of the confusion here. 1. **Across workers (sharding).** Each of N machines must see a different part of the dataset, or you are training on N copies of the same data per epoch. 2. **Across replicas within a worker (batch splitting).** Each global batch is sliced along the batch axis and one slice goes to each local GPU. The second is automatic and uninteresting. The first is auto-sharding, and it is where the failure modes live. ## The policies `tf.data.Options().experimental_distribute.auto_shard_policy` takes a value from `tf.data.experimental.AutoShardPolicy`: - **`AUTO`** (default) — try `FILE`; fall back to `DATA` if that is not possible. - **`FILE`** — split the *input files* across workers. Worker 0 reads files 0, N, 2N…; worker 1 reads 1, N+1… Each record is read exactly once cluster-wide. This is the efficient path. - **`DATA`** — every worker reads the *entire* dataset and drops every record whose index does not belong to it. Correct, and horribly wasteful: with 8 workers you do 8x the read I/O and 8x the parsing to consume the data once. - **`OFF`** — no sharding. Every worker trains on the whole dataset, which for synchronous data parallelism means every replica sees the same examples and you have effectively multiplied the batch by nothing useful. Only correct when you shard yourself upstream. Attach the option with `dataset.with_options(options)`. ## When FILE sharding is unavailable `FILE` needs a file-based source (`TFRecordDataset`, `list_files`, and similar) and at least as many files as workers. Both conditions fail routinely: - The pipeline starts from `from_tensor_slices` on an in-memory array — no files at all. - The data lives in one giant TFRecord. One file, eight workers. - The dataset is generated by `from_generator`, which cannot be split by file. With `AUTO` these fall back to `DATA` and the job trains correctly but slowly, and the reason is buried in a log line. With an explicit `FILE` policy, TensorFlow raises rather than degrading — which is why setting it explicitly is a good production discipline: you would rather fail at start than discover a 5x input bottleneck a week later. The design implication is direct: **shard your data into many more files than you will ever have workers** — a few hundred shards of tens of megabytes each is the conventional shape. This is a data-preparation decision made long before the training job runs. ## Skew `FILE` sharding assigns whole files, so if the files differ in record count, some workers finish their epoch earlier than others. Under synchronous training the fastest workers then wait, and the job runs at the pace of the worker with the most records. Roughly equal shard sizes matter more than shard count alone. ## Shuffling and determinism Shuffle placement interacts with sharding. If you shuffle a file list before sharding and the shuffle is seeded differently per worker, workers can disagree about which files exist where and the same record may be read twice or not at all. Keep the file listing deterministic across workers, and do record-level shuffling *after* sharding with a shuffle buffer. `tf.data.Dataset.list_files` shuffles by default, so passing `shuffle=False` and shuffling explicitly is the safer construction under multi-worker. ## Taking control explicitly `strategy.distribute_datasets_from_function(fn)` calls `fn` once per input pipeline with a `tf.distribute.InputContext`. Useful members: `num_input_pipelines`, `input_pipeline_id`, and `get_per_replica_batch_size(global_batch_size)`. Note the last one — with this API you batch with the **per-replica** size, because you are constructing one pipeline per replica-group rather than handing over a global-batched dataset. This is the escape hatch when your data lives somewhere auto-sharding cannot reason about: a sharded database query, a custom object store, a partition scheme keyed on something other than files. You call `dataset.shard(num_shards, index)` yourself, or better, push the shard predicate into the query so you never read the discarded rows. ## Verifying it worked The cheap check: count the records one worker consumes in an epoch and confirm it is roughly the total divided by the worker count. If every worker consumes the whole dataset, sharding is `OFF` or the pipeline is not the one you distributed. This is worth doing once per new data layout, because both failure modes — reading everything, and reading everything and throwing most of it away — produce correct-looking loss curves.

  • Why does distribute_datasets_from_function batch with the per-replica size rather than the global one?
    Because `fn` builds one pipeline per input pipeline rather than a single global dataset the strategy then splits. `InputContext.get_per_replica_batch_size(global)` does the division for you, so the batches the pipeline emits are already replica-sized. Passing the global size here silently multiplies your effective batch.
  • You have one 200 GB TFRecord file and eight workers. What do you do?
    Re-shard the data into many files — commonly a few hundred, roughly equal in record count — as a preparation step. With one file, AUTO falls back to DATA sharding and all eight workers read and parse the whole 200 GB to consume it once. No runtime option fixes a single-file layout; the fix is upstream.
  • What breaks if shard sizes are uneven under FILE sharding?
    Workers get whole files, so one worker may hold noticeably more records than another. Synchronous training advances at the pace of the slowest worker, so the imbalance shows up as idle time on the fast machines and an epoch that takes as long as the heaviest shard. Aim for roughly equal record counts per file.
  • How do you confirm sharding actually happened?
    Count the records one worker consumes in an epoch and compare with the dataset total divided by the worker count. If a worker consumes everything, the policy is OFF, the pipeline you distributed is not the one being iterated, or the fallback is doing whole-dataset reads. Both wrong configurations still produce plausible loss curves.

FILE sharding is handing each reader a different chapter; DATA sharding is handing everyone the whole book and telling them to read only every eighth page.

saying these in an interview costs you the question

  • Assuming every worker automatically gets distinct data
  • Training multi-worker from one giant TFRecord file
  • Leaving auto_shard_policy at AUTO and never checking the fallback
  • Sharding after shuffling the file list per worker
  • Batching with the global size inside a dataset function

context