How does switching a loss from a batch mean to a plain sum change every gradient in the graph?
answer
- a reduction has a local gradient
- 1/N versus 1 handed to each element
- direction survives, magnitude does not
- the learning rate absorbs the factor N
basics
~20 sEvery gradient is multiplied by the batch size N. A mean reduction hands each element a local gradient of 1/N; a plain sum hands it 1. The descent direction is identical, the step is N times larger.
solid answer
~50 sA mean over N per-example losses is `L = (1/N) * sum(l_i)`, so the local gradient the reduction hands to each `l_i` is `1/N`. A plain sum hands each one `1`. That factor multiplies straight through the rest of the reverse sweep, so *every* parameter gradient in the graph is exactly N times larger under a sum. The direction is unchanged -- it is one positive scalar -- but the effective step size is not, so a learning rate tuned under a mean will typically diverge at a batch size of 256 under a sum. Everything calibrated in gradient units moves with it: clipping thresholds, the ratio between an L2 penalty and the data term, any warning threshold on gradient norm. That coupling is why the mean is the usual default -- it keeps hyperparameters roughly invariant to batch size.
go deeper
Know that a mean divides the summed loss by the number of examples, and that this divisor reaches the gradients too rather than only the number printed in the log.
Be able to state the local gradient of each reduction -- 1/N for a mean, 1 for a sum -- and explain why a factor applied at the loss multiplies every parameter gradient in the graph.
Diagnose it from symptoms: divergence right after a change that supposedly only touched the loss, a loss magnitude near N times the per-example scale, or a gradient norm that jumped by exactly the batch size. Know that clipping thresholds and regularisation balance move with the scale.
Set the convention and make it explicit -- what the denominator is, what happens under gradient accumulation and masking, and how learning rate is coupled to batch size -- so that throughput changes do not silently retune the model. Treat any reduction convention that is implicit as a latent hyperparameter.
## The local gradient of a reduction A reduction node takes many numbers and returns one. Two of them dominate loss code: ``` L_sum = l_1 + l_2 + ... + l_N -> dL/dl_i = 1 L_mean = (l_1 + ... + l_N) / N -> dL/dl_i = 1/N ``` That is the entire mechanical difference: the reduction hands every per-example loss either `1` or `1/N` as its local gradient. Everything downstream in the reverse sweep is multiplied by whatever the reduction handed out, so a single scalar factor of `N` propagates to every weight gradient, every bias gradient, and every intermediate gradient in the graph. ## What changes and what does not **Unchanged: the direction.** Multiplying a gradient vector by a positive scalar does not rotate it. The sum-reduced and mean-reduced losses have gradients that are exact positive multiples of each other, so plain gradient descent takes a step along the same ray. **Changed: the magnitude, and everything calibrated against it.** - *Effective learning rate.* With a fixed learning rate, a sum reduction takes a step N times longer. At a batch size of 256 that is a 256-fold change, which usually means immediate divergence to non-finite values within a handful of steps. - *Batch-size coupling.* Under a mean, doubling the batch leaves gradient magnitudes roughly where they were (the average of twice as many similar terms), so the learning rate stays roughly valid. Under a sum, doubling the batch doubles every gradient, so the learning rate must be halved to keep the same step. This invariance is the practical reason the mean is the default. - *Gradient clipping.* A norm threshold is in gradient units. Under a sum, the same threshold clips N times more aggressively and can silently turn training into normalised, direction-only steps. - *Regularisation balance.* A weight penalty added to the loss keeps its own scale while the data term is multiplied by N, so the effective regularisation strength drops by a factor of N relative to the data term. The model looks under-regularised for reasons nothing in the regulariser changed. - *Adaptive optimisers hide it, partially.* An optimiser that divides an estimate of the gradient by an estimate of its own magnitude is approximately invariant to a global rescale once its moment estimates have adapted, so the divergence may be delayed or muted rather than immediate. It is not fully invariant -- the early steps before the estimates settle, and any small constant added to the denominator for numerical safety, both feel the scale. Do not rely on the optimiser to absorb the change. ## Where the choice genuinely matters The mean is not automatically right; the real question is *the mean over what*. - **Variable-length data.** If a batch of sequences is padded to a common length and masked, dividing by the padded element count instead of the count of valid elements makes every gradient depend on how much padding a batch happened to contain. Two batches with identical content but different padding then produce different-sized updates. The denominator must be the number of valid elements. - **Gradient accumulation.** When a large batch is processed as several micro-batches whose gradients accumulate before one update, each micro-batch loss must be divided by the *total* number of examples, not by the micro-batch size. Dividing by the micro-batch size makes the accumulated gradient a sum of averages, which is the number of micro-batches times too large -- exactly the mean-versus-sum bug wearing a different hat. - **Multi-term losses.** When several loss terms are added with weights, the reduction of each term must be consistent. Mixing a summed term with an averaged term makes the effective weights depend on batch size, which is a hyperparameter that silently changes every time someone tunes throughput. ## Diagnosing it in the wild The signature of an accidental switch is abrupt and scale-shaped: the loss becomes non-finite in a few steps after a change that "only touched the loss reduction", or the reported loss value jumps by roughly the batch size while the model behaves normally for a while under an adaptive optimiser. Two quick checks settle it. First, compare the reported loss magnitude to the expected per-example scale; a loss that is roughly N times a plausible per-example value is a summed loss. Second, compare the global gradient norm before and after the change; a clean factor of the batch size is conclusive, since a genuine modelling change essentially never produces exactly that ratio. The rule to carry away: a reduction is not a formatting choice about how the loss is displayed. It is a node with a local gradient, and its local gradient scales the entire graph.
- When a large batch is split into micro-batches whose gradients accumulate, what must the divisor be?The total number of examples across all micro-batches, not the micro-batch size. Dividing each micro-batch loss by its own size and then adding the gradients produces a sum of averages, which overshoots by the number of micro-batches. Either scale each micro-batch loss by its share of the total, or accumulate sums and divide once before the update.
- Does an adaptive optimiser that normalises by a running estimate of gradient magnitude make this switch harmless?Only approximately, and only after its estimates have adapted. Such an optimiser divides by a quantity that scales with the gradient, so a global rescale largely cancels in steady state. It does not cancel during the first steps while the estimates are still catching up, and the small constant added to the denominator for numerical safety does not scale, so very small gradients are affected differently. Treat the invariance as a cushion, not a guarantee.
- A batch of padded sequences is averaged over all positions including padding. Why is that a bug?The divisor then depends on how much padding the batch happened to contain, so two batches with identical real content but different padding produce different gradient magnitudes. Effectively the per-example weight varies with batch composition. Averaging over the count of valid elements restores a stable scale and makes updates independent of how sequences were bucketed.
saying these in an interview costs you the question
- Says the gradient direction changes, not just the magnitude
- Thinks the reduction affects only the displayed loss value
- Assumes an adaptive optimiser makes the scale fully irrelevant
- Divides an accumulated micro-batch loss by the micro-batch size
- Averages over padded positions rather than valid elements