Under tf.distribute MirroredStrategy, is model.fit's batch_size global or per GPU?
answer
- The number you write is not per device
- Replicas split one batch, they do not each get one
- num_replicas_in_sync is the multiplier
- Four GPUs, batch 256, sixty-four each
- Derive global from per-replica, not the reverse
basics
~10 sGlobal. Keras splits each batch evenly across the replicas, so batch_size=256 on four GPUs means 64 examples per GPU per step. To keep per-GPU work constant, raise the global number as you add devices.
solid answer
~40 sIn `tf.distribute`, every batch size you write is the **global** batch size — the number of examples consumed per training step across all replicas together. With `MirroredStrategy` on four GPUs and `model.fit(..., batch_size=256)`, each replica gets a 64-example slice, computes its own forward and backward pass, and the gradients are combined before the weights update. The same rule applies to `tf.data`: you call `dataset.batch(global_batch_size)` and hand the batched dataset to `fit`, not a per-replica batch. The practical consequence is that moving an existing script from one GPU to four without touching `batch_size` does not speed up each step much — each device now does a quarter of the work. The usual move is `per_replica_batch * strategy.num_replicas_in_sync`, so per-device work stays where you tuned it.
code
python · 13 linesimport tensorflow as tf
strategy = tf.distribute.MirroredStrategy()
per_replica_batch_size = 64
global_batch_size = per_replica_batch_size * strategy.num_replicas_in_sync
with strategy.scope():
model = tf.keras.Sequential([tf.keras.layers.Dense(1)])
model.compile(optimizer="adam", loss="mse")
x = tf.random.normal((1024, 8))
y = tf.random.normal((1024, 1))
model.fit(x, y, batch_size=global_batch_size, epochs=1)go deeper
Remember one sentence: in tf.distribute the batch size is global, and the replicas split it. Be ready to say that four GPUs with batch_size=256 means 64 examples per GPU.
Explain the mechanics: each replica computes gradients on its slice, an all-reduce combines them, and every mirrored copy applies the same update. Show that you derive the global batch as per-replica times num_replicas_in_sync.
Be ready to explain why a naive port to multi-GPU gains almost nothing, and to reason about small per-device batches, all-reduce overhead, and keeping the global batch an exact multiple of the replica count.
Own the tradeoff between scaling the global batch for throughput and the training-dynamics cost of doing so, and set a team convention so scripts declare a per-replica batch and derive the global one rather than hardcoding device counts.
## The one rule Everywhere in `tf.distribute`, a batch size is a **global** batch size: the number of examples the whole replica group consumes in one training step. This is a deliberate API choice, and it is the single most common source of confusion when a script moves from one device to several. A *replica* is one copy of the model running on one device. `MirroredStrategy` creates one replica per visible GPU on a single machine, and `strategy.num_replicas_in_sync` tells you how many there are (it is 1 when you run on CPU only, so the same script works unchanged on a laptop). ## What actually happens to a batch When you call `model.fit(x, y, batch_size=256)` inside a strategy, Keras builds a distributed dataset. Each step, that dataset produces one global batch of 256 examples and splits it along the first axis into `num_replicas_in_sync` pieces. On four GPUs each replica receives 64 examples. Every replica runs the forward pass and computes gradients on its own slice, using its own copy of the weights — the copies are kept identical, which is what "mirrored" means. The per-replica gradients are then combined across devices with an all-reduce, and the single combined gradient updates every copy of every variable. Because all replicas apply the same update, the copies never drift apart. Mathematically the result is the same as one device processing all 256 examples in one batch (given a correctly scaled loss). That equivalence is the whole point of synchronous data parallelism, and it is why the API defines batch size globally: the number you write is the number the *optimization* sees, regardless of how many devices you happen to own. ## The surprise A person who tuned `batch_size=64` on one GPU and then wraps the model in `MirroredStrategy` on four GPUs often reports that training got barely faster. That is expected: each GPU is now doing 16 examples per step. Small per-device batches under-utilize the GPU, and the fixed cost of the all-reduce is now amortized over less work, so scaling efficiency is poor. The standard fix is to define the per-replica batch size explicitly and derive the global one: per_replica = 64; global = per_replica * strategy.num_replicas_in_sync Now each GPU still sees 64 examples, and adding devices increases throughput rather than splitting the existing work. Note that this raises the effective batch the optimizer sees, which is a training-dynamics decision, not an API one — the point here is that `tf.distribute` will not make that decision for you. ## tf.data and steps If you pass a `tf.data.Dataset` instead of arrays, you must not also pass `batch_size` — the dataset is already batched, and it must be batched with the **global** size. Likewise `steps_per_epoch` counts global steps: total examples divided by global batch size, not by per-replica batch size. Getting this wrong is a silent bug — the run completes and simply sees a different amount of data than you intended. ## Divisibility and partial batches If the global batch does not divide evenly by the replica count, the split is uneven and some replica does more work than the others, which wastes time on every step since synchronous training moves at the pace of the slowest replica. Keeping the global batch an exact multiple of `num_replicas_in_sync` is the simple discipline. The final batch of an epoch is often smaller than the rest; `tf.distribute` handles such partial batches, but if you write a custom loop rather than using `fit`, a partial batch is exactly the case where a hand-rolled loss average over the *actual* tensor length silently disagrees with the global-batch scaling the framework assumes. ## What this does not change Global batch semantics affect how data is fed and how loss is scaled. They do not change your model code, your layers, or your metrics, and they are unrelated to how many machines are involved — a multi-worker job uses the same rule, with the global batch spread across every replica on every worker.
- And if you pass a tf.data.Dataset to fit instead of arrays, where does the batch size go?Into the dataset: you call `dataset.batch(global_batch_size)` yourself and pass no `batch_size` to `fit` — passing both raises an error. The strategy then splits each already-formed global batch across replicas. `steps_per_epoch` must also be counted in global steps.
- What does strategy.num_replicas_in_sync return when no GPU is visible?1. `MirroredStrategy` falls back to a single CPU replica, so the same script runs unchanged on a machine with no accelerator. That is why deriving the global batch from `num_replicas_in_sync` is safer than hardcoding a device count — the code is correct on a laptop and on an eight-GPU host.
- Does splitting a batch across four GPUs change the result compared with one GPU?With a correctly scaled loss, no — synchronous data parallelism reproduces the single-device update for the same global batch, because gradients are combined before the weights change. Numerics can differ slightly from reduction order, and any per-batch statistic computed locally, such as batch normalization, is computed per replica rather than over the global batch.
Think of a global batch as a stack of paperwork handed to a team, not to each person: adding people makes the stack finish sooner, it does not give everyone a full stack.
saying these in an interview costs you the question
- Thinking batch_size is multiplied by the number of GPUs
- Passing both a batched dataset and batch_size to fit
- Expecting a 4x speedup without raising the global batch
- Counting steps_per_epoch with the per-replica batch size
- Assuming per-device batch shrinks only on the last step