skip to content

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

level: seniorimportance: must knowfreq 55%

answer

  1. Gradients are summed, not averaged
  2. Find the denominator of the loss
  3. Divide by the global batch, always
  4. reduction="none" first, then scale
  5. Regularization terms need their own division

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.

solid answer

~40 s

Under `tf.distribute`, the gradients each replica computes are **summed** across replicas before they are applied. So each replica must contribute its share of a global average, not a local average. Concretely: construct the loss object with `reduction="none"` so it returns a per-example vector, then divide the sum of that vector by the **global** batch size — `tf.nn.compute_average_loss(per_example_loss, global_batch_size=GLOBAL_BATCH_SIZE)` does exactly this. If you instead write `tf.reduce_mean(per_example_loss)`, each replica divides by its own slice size, and after the sum the effective gradient is `num_replicas_in_sync` times too large — training does not crash, it just behaves as though you silently multiplied the learning rate. For reporting, `strategy.run` returns a per-replica value, which you collapse with `strategy.reduce(tf.distribute.ReduceOp.SUM, per_replica_losses, axis=None)` — SUM, not MEAN, because each replica already carries its global-scaled fraction.

code

python · 25 lines
python
import tensorflow as tf

strategy = tf.distribute.MirroredStrategy()
GLOBAL_BATCH_SIZE = 64

with strategy.scope():
    model = tf.keras.Sequential([tf.keras.layers.Dense(1)])
    optimizer = tf.keras.optimizers.SGD(0.01)
    loss_fn = tf.keras.losses.MeanSquaredError(reduction="none")

def step_fn(inputs):
    x, y = inputs
    with tf.GradientTape() as tape:
        per_example_loss = loss_fn(y, model(x, training=True))
        loss = tf.nn.compute_average_loss(
            per_example_loss, global_batch_size=GLOBAL_BATCH_SIZE
        )
    grads = tape.gradient(loss, model.trainable_variables)
    optimizer.apply_gradients(zip(grads, model.trainable_variables))
    return loss

@tf.function
def train_step(dist_inputs):
    per_replica = strategy.run(step_fn, args=(dist_inputs,))
    return strategy.reduce(tf.distribute.ReduceOp.SUM, per_replica, axis=None)

go deeper

for a junior

Know that model.fit handles this for you, and that hand-written distributed loops must scale the loss by the global batch size rather than calling reduce_mean.

for a middle

Explain why: gradients are summed across replicas, so a per-replica mean inflates the update by the replica count. Show the reduction="none" plus compute_average_loss pattern from memory.

for a senior

Diagnose it in the wild — a run that is stable on two GPUs and diverges on eight with identical hyperparameters — and cover the second-order cases: regularization terms, partial final batches, and metrics that need no scaling.

for a principal

Make this un-reproducible at the org level: prefer fit or a shared step-function helper over hand-rolled loops, and require a two-replica-versus-one-replica equivalence check in CI so a factor-of-R bug cannot ship.

## Why this exists at all `model.fit` handles loss scaling for you. The moment you write a custom loop — `GradientTape`, `strategy.run`, `apply_gradients` — you take ownership of it, and this is the single most common correctness bug in distributed TensorFlow code. It produces no exception and no warning; it produces a number that is wrong by an exact integer factor. ## The mechanism Each replica runs the step function on its slice of the global batch and computes gradients from its own loss. When the optimizer applies those gradients under a strategy, they are combined with an all-reduce **sum**, and the summed gradient updates every mirrored copy of every variable. The target is to reproduce what a single device would compute on the full global batch, which is the mean over all global-batch examples: `L_global = (1 / B_global) * sum over all examples of l_i` Since gradients add, replica *r* should contribute `(1 / B_global) * sum over its own examples`. That means dividing by `B_global` — the global batch — even though the replica only holds `B_global / R` examples. Write `tf.reduce_mean(per_example_loss)` instead and replica *r* contributes `(R / B_global) * sum over its own examples`. Add up `R` replicas and the total is `R` times the intended gradient. Everything still runs; the model just trains as if the learning rate were multiplied by the number of GPUs. Symptoms are diffuse: divergence on 8 GPUs but not on 2, or a run that is stable single-device and unstable multi-device with "the same" hyperparameters. ## The correct code shape Three pieces, all required: 1. **`reduction="none"` on the loss object.** Keras loss classes reduce internally by default (over the batch). You need the raw per-example vector so *you* control the denominator. In Keras 3, the reduction is a string argument on the loss constructor. 2. **`tf.nn.compute_average_loss(per_example_loss, global_batch_size=GLOBAL_BATCH_SIZE)`.** It sums the per-example losses and divides by the global batch size. It also accepts `sample_weight`. Writing `tf.reduce_sum(per_example_loss) / GLOBAL_BATCH_SIZE` by hand is equivalent; the helper exists so the intent is legible in review. 3. **`strategy.reduce(tf.distribute.ReduceOp.SUM, per_replica_losses, axis=None)`** to turn the `PerReplica` value that `strategy.run` returns into a scalar for logging. SUM is right because each replica's value is already its fraction of the global mean; using MEAN here would under-report by `R`. ## Regularization and other extra losses Weight-decay and activity-regularization terms in `model.losses` are *not* per-example — there is one such term per layer, and every replica computes the same value. If you just add them, the sum across replicas multiplies them by `R`. `tf.nn.scale_regularization_loss` divides such a term by the number of replicas so the total comes out right. Skipping this is the second-order version of the same bug, and it is why blindly adding `sum(model.losses)` into a distributed step is wrong. ## Partial batches The last batch of an epoch is often shorter, and with an uneven split some replica may even receive zero examples. This is exactly where dividing by `tf.shape(per_example_loss)[0]` — a tempting "dynamic" version of the mean — goes wrong: the denominators no longer add up to the global batch, and a replica with an empty slice can produce `nan` from a zero division. Dividing by the fixed `GLOBAL_BATCH_SIZE` constant is well-defined in every case. The resulting last step is weighted slightly differently from a single-device run, which is the accepted tradeoff. ## Metrics behave differently Stateful Keras metrics accumulate their own variables, and under a strategy those accumulators are mirrored and aggregated when you read `.result()`. So metrics do **not** need the global-batch treatment — updating a metric inside the replica function and reading it in cross-replica context gives the right answer. Confusing the loss rule with the metric rule leads people to divide metrics by the replica count and under-report accuracy. ## Mixed precision interaction If you use a loss-scaling optimizer for float16, the loss scaling and the batch scaling are independent: scale the loss for numeric range *after* averaging over the global batch, and unscale the gradients before applying. Both multiplications commute, but forgetting one of them is again a silent factor error. ## The review heuristic In any distributed custom loop, find the denominator of the loss. If it is a per-replica batch size, a dynamic `tf.shape`, or an implicit `reduce_mean`, it is a bug. If it is the global batch constant, it is right.

  • What happens to model.losses — the regularization terms — in that step function?
    They are per-layer, not per-example, so every replica computes the identical value and the cross-replica sum multiplies it by the replica count. Divide each such term by `strategy.num_replicas_in_sync` before adding it, which is what `tf.nn.scale_regularization_loss` does. Otherwise your effective weight decay scales with your GPU count.
  • Why strategy.reduce with ReduceOp.SUM rather than MEAN when logging the loss?
    Because each replica already divided by the global batch size, so its value is a fraction of the global mean, not the mean itself. Summing those fractions gives the global mean. Using MEAN would divide once more and under-report the loss by the number of replicas.
  • Why not divide by the actual number of examples the replica received, using tf.shape?
    Because the denominators then no longer sum to the global batch, so partial final batches get a different effective weight than you intended, and a replica handed an empty slice divides by zero and yields nan. A fixed global constant is well-defined for every step and matches what the framework assumes.
  • Do stateful Keras metrics need the same global-batch scaling?
    No. Metric variables are mirrored and aggregated when you read `result()`, so updating a metric inside the replica function and reading it in cross-replica context already reports the correct global value. Applying the loss rule to metrics is a common over-correction that under-reports accuracy by the replica count.

Each replica is filling in one part of a shared average: it must divide by the size of the whole class, not by the size of its own row.

saying these in an interview costs you the question

  • Calling tf.reduce_mean on the per-replica loss
  • Leaving the loss object's default reduction in place
  • Adding sum(model.losses) without dividing by replica count
  • Using strategy.reduce with MEAN after global scaling
  • Dividing by the dynamic batch length from tf.shape

context