skip to content

TensorFlow

TensorFlow covers the same ground as PyTorch with different defaults: tf.function graphs, tf.data input pipelines, Keras as the model API, tf.distribute for scale, and SavedModel plus TF Serving for deployment. Common wherever a production stack was built a few years ago.

on this pageshow

explore

questions

page 1 of 2

Under tf.distribute MirroredStrategy, is model.fit's batch_size global or per GPU?

level: juniorimportance: must knowfreq 70%

answer

  1. The number you write is not per device
  2. Replicas split one batch, they do not each get one
  3. num_replicas_in_sync is the multiplier
  4. Four GPUs, batch 256, sixty-four each
  5. Derive global from per-replica, not the reverse

basics

~10 s

Global. 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 s

In `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 lines
python
import 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

for a junior

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.

for a middle

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.

for a senior

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.

for a principal

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

context

open as a page

How do you send a prediction request to TensorFlow Serving's REST API?

level: juniorimportance: must knowfreq 58%

basics

~10 s

POST JSON to http://host:8501/v1/models/MODEL:predict with an "instances" list of examples; the response is a JSON object with a "predictions" list in the same order. Port 8501 is REST; 8500 is gRPC.

open as a page

Why does tf.data.Dataset.prefetch(tf.data.AUTOTUNE) go last in a pipeline?

level: middleimportance: must knowfreq 72%

basics

~20 s

prefetch decouples producing elements from consuming them: a background thread fills a buffer while the accelerator trains on the previous batch. Placed last, after batch, its buffer holds ready-to-use batches, so step N+1's input work overlaps step N's compute.

open as a page

In tf.data, what does shuffle(buffer_size) actually do, and why does its position matter?

level: middleimportance: must knowfreq 78%

basics

~20 s

tf.data's shuffle keeps a buffer of buffer_size elements, emits one at random from it, and refills the slot from the source. It is a sliding-window shuffle, not a global one, so a small buffer over sorted data barely mixes anything.

open as a page

What must be created inside a tf.distribute strategy.scope(), and why?

level: middleimportance: must knowfreq 80%

basics

~20 s

Anything that creates variables: the model, the optimizer, and any metric or custom tf.Variable. Inside the scope those become mirrored variables with one synchronized copy per replica. Datasets, fit calls, and the training step itself go outside.

open as a page

In TensorFlow 2.21, what does tf.keras resolve to, and how do you get Keras 2 back?

level: middleimportance: must knowfreq 55%

basics

~20 s

tf.keras in TensorFlow 2.21 is Keras 3, not the Keras 2 that older TensorFlow bundled. To get Keras 2 semantics back you install the separate tf_keras package and set the environment variable TF_USE_LEGACY_KERAS=1 before importing TensorFlow.

open as a page

What does a TensorFlow SavedModel directory contain, and why does TF Serving need it?

level: middleimportance: must knowfreq 72%

basics

~20 s

A SavedModel is a directory: saved_model.pb holds the serialized graph plus signature definitions, variables/ holds the weight checkpoint, assets/ holds files the graph reads, and fingerprint.pb identifies the export. TF Serving needs the graph, not just weights.

open as a page

How do you set and inspect the serving_default signature of a SavedModel?

level: middleimportance: must knowfreq 64%

basics

~20 s

Pass signatures= to tf.saved_model.save() mapping a key such as serving_default to a concrete function, or let Keras 3's model.export() create it. Inspect with saved_model_cli show --dir m/1 --all, or load the model and read .signatures.

open as a page

Why does a tf.function retrace, and how do you stop it retracing on every call?

level: middleimportance: must knowfreq 70%

basics

~20 s

A tf.function caches one graph per input signature — tensor dtypes and shapes, plus the value of any Python argument. New shapes or changing Python arguments therefore force a fresh trace. Pass tensors instead of Python scalars, or pin an input_signature with None dimensions.

open as a page

What does tf.function change about how TensorFlow executes your Python code?

level: middleimportance: must knowfreq 82%

basics

~20 s

tf.function traces the decorated Python function once per input signature, recording the TensorFlow ops it calls into a dataflow graph, then runs that cached graph on later calls. Python executes only at trace time; the graph executes on every call.

open as a page

How does a tf.GradientTape training step compute and apply gradients in TensorFlow?

level: middleimportance: must knowfreq 76%

basics

~10 s

tf.GradientTape records the forward pass that runs inside its with block. tape.gradient(loss, model.trainable_variables) replays that recording backwards to produce gradients, and optimizer.apply_gradients(zip(grads, variables)) writes the update into the variables.

open as a page

A tf.data pipeline reading thousands of TFRecord shards starves the GPU — how do you fix it?

level: seniorimportance: must knowfreq 60%

basics

~10 s

Confirm the job is input-bound, then add parallelism where the time goes: read shards concurrently with interleave(num_parallel_calls=AUTOTUNE) or TFRecordDataset(num_parallel_reads=...), parse vectorized after batch with tf.io.parse_example, map with num_parallel_calls, and end with prefetch.

open as a page

In a custom TensorFlow training loop under tf.distribute, how do you reduce the loss?

level: seniorimportance: must knowfreq 55%

basics

~10 s

Scale by the global batch size, not the per-replica one. Build the loss with reduction="none", then call tf.nn.compute_average_loss(per_example_loss, global_batch_size=GLOBAL_BATCH). Gradients are summed across replicas, so a per-replica mean inflates them by the replica count.

open as a page

How does TF Serving pick which SavedModel version to serve, and how do you roll one out?

level: seniorimportance: must knowfreq 52%

basics

~20 s

Versions are numeric subdirectories under the model base path. TF Serving polls that path and by default serves the highest number, loading the new version fully before switching traffic and then unloading the old one. A model config file changes the policy and adds labels.

open as a page

In tf.data, how do Dataset.from_tensors and Dataset.from_tensor_slices differ?

level: juniorimportance: should knowfreq 55%

basics

~10 s

Dataset.from_tensors makes a one-element dataset holding the whole tensor as a single item. Dataset.from_tensor_slices splits along the first axis, yielding one element per row. Training on N examples almost always wants from_tensor_slices.

open as a page

In TensorFlow, why does adding a float32 tensor to a float64 tensor raise an error?

level: juniorimportance: should knowfreq 58%

basics

~20 s

TensorFlow ops require both operands to carry the same dtype and do not promote silently the way NumPy does, so mixing float32 and float64 raises InvalidArgumentError. Convert one side explicitly with tf.cast before combining them.

open as a page

What does from_logits=True do in tf.keras.losses.SparseCategoricalCrossentropy?

level: juniorimportance: should knowfreq 56%

basics

~20 s

from_logits=True tells the loss that the model's outputs are raw, unnormalized scores, so the loss applies softmax internally in a numerically stable way. The default, False, means the loss assumes the values are already probabilities.

open as a page

Why does random augmentation inside tf.data Dataset.map() sometimes repeat identically?

level: middleimportance: should knowfreq 58%

basics

~20 s

Dataset.map traces its function once into a graph rather than running it per element. Plain Python calls such as random.random() or numpy.random execute only at trace time, so their result is baked in as a constant and every element gets the same value.

open as a page

How does TensorFlow's MultiWorkerMirroredStrategy find its peer workers?

level: middleimportance: should knowfreq 52%

basics

~20 s

Through the TF_CONFIG environment variable. It is a JSON string holding a "cluster" map of every worker's host:port and a "task" entry giving this process its own type and index. Every worker runs the same script with a different index.

open as a page

Why does a raw tf op fail on the output of keras.Input in Keras 3?

level: middleimportance: should knowfreq 42%

basics

~20 s

keras.Input returns a symbolic KerasTensor placeholder, not a tf.Tensor. Keras 3 is backend-agnostic, so handing that placeholder to a TensorFlow op raises an error. Use keras.ops, or put the TensorFlow op inside a layer's call().

open as a page

Why does a Python print() inside a tf.function only run on the first call?

level: middleimportance: should knowfreq 60%

basics

~20 s

Plain Python statements execute only while TensorFlow traces the function body into a graph; after that, calls run the graph, which contains TensorFlow ops and nothing Python. Use tf.print, which becomes an actual graph node, to print on every call.

open as a page

How do you attach a learning-rate schedule to a TensorFlow Keras optimizer?

level: middleimportance: should knowfreq 48%

basics

~20 s

Pass a LearningRateSchedule object as the optimizer's learning_rate argument instead of a float. The optimizer evaluates it against its own step counter on every update, so decay_steps and similar arguments count optimizer steps — batches — never epochs.

open as a page

Why does tf.GradientTape.gradient() return None in TensorFlow, and how do you fix it?

level: middleimportance: should knowfreq 52%

basics

~20 s

None means the tape found no recorded path from the loss back to that source. Usual causes: the source is not a watched variable, the computation happened outside the tape's context, or the path runs through a non-differentiable op. Fix the connection, not the symptom.

open as a page

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

level: seniorimportance: should knowfreq 50%

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.

open as a page

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

level: seniorimportance: should knowfreq 45%

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.

open as a page

Your Keras fit() on a repeated tf.data.Dataset never ends an epoch and ignores shuffle=True — why?

level: seniorimportance: should knowfreq 50%

basics

~20 s

A 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.

open as a page

Why does a SavedModel exported without an input_signature reject a different batch size?

level: seniorimportance: should knowfreq 46%

basics

~20 s

The signature written into a SavedModel is a concrete function, and a concrete function fixes the exact dtype and shape it was traced with. Trace with tf.TensorSpec([None, ...]) so the exported signature accepts any batch size.

open as a page

How does AutoGraph handle a Python if statement inside a tf.function?

level: seniorimportance: should knowfreq 52%

basics

~20 s

AutoGraph source-rewrites the function before tracing. If the condition is a tensor, the if becomes a tf.cond node and both branches are traced; if the condition is a plain Python value, the interpreter decides at trace time and only the taken branch enters the graph.

open as a page

How do you enable mixed_float16 in TensorFlow, and what must you fix afterwards?

level: seniorimportance: should knowfreq 46%

basics

~10 s

Call tf.keras.mixed_precision.set_global_policy("mixed_float16"): layers then compute in float16 while keeping float32 weights. Two fixes follow — force the output layer to float32, and apply loss scaling so small float16 gradients do not underflow to zero.

open as a page

Your TensorFlow training loss turns NaN after a few hundred steps — how do you diagnose it?

level: seniorimportance: should knowfreq 44%

basics

~20 s

Find the first op that produces a non-finite value rather than guessing. tf.debugging.enable_check_numerics() raises at that op with a stack trace; then check inputs for NaN, watch the gradient global norm for a spike, and only then reach for clipping or a lower rate.

open as a page

showing 1–30 of 35