In gradient accumulation, why is each micro-batch's loss divided by the number of micro-batches?
answer
- Mean of means is not the mean
- Which reduction: sum or mean?
- Divide by micro-batches, or by samples
- A constant factor is a learning-rate change
- Short final group needs its own divisor
basics
~20 sBecause each micro-batch's loss is usually already averaged over its own samples, so adding N of them gives N times the true full-batch gradient. The extra 1/N factor restores the mean over all samples and keeps the step size honest.
solid answer
~50 sThe accumulated gradient has to equal the gradient of the average loss over all the samples in the group. Which divisor achieves that depends on how the per-micro-batch loss is reduced. If it is a **mean** over its own `m` samples, then adding `N` of those means gives `N` times the group mean, so each micro-batch loss must be divided by `N` before its backward pass. If it is a **sum** over its `m` samples, the accumulated total must be divided by the total sample count `N*m`. Getting this wrong throws nothing: the run keeps going with a gradient that is a constant multiple of the intended one, which is exactly equivalent to silently multiplying the learning rate. The classic version is eight micro-batches of four summed and divided by eight instead of thirty-two — a step four times too large, later blamed on the learning rate.
code
python · 19 lines# 8 samples = 4 micro-batches of 2; per-sample gradient of 0.5*(w*x - y)**2 is (w*x - y)*x
data = [(1.0, 2.0), (2.0, 3.0), (3.0, 7.0), (4.0, 8.0),
(1.5, 4.0), (2.5, 5.0), (3.5, 6.0), (0.5, 1.0)]
w, micro = 0.5, 2
n_micro = len(data) // micro
def grads(batch):
return [(w * x - y) * x for x, y in batch]
full_batch = sum(grads(data)) / len(data)
right = wrong = 0.0
for i in range(0, len(data), micro):
chunk = data[i:i + micro]
micro_mean = sum(grads(chunk)) / len(chunk) # loss averaged inside the micro-batch
wrong += micro_mean # bug: never divided by the 4 micro-batches
right += micro_mean / n_micro # correct: the extra 1/N
print(full_batch, right, wrong) # -9.4375 -9.4375 -37.75go deeper
Remember that the loss must be scaled down before each micro-batch's backward pass, and that forgetting it makes the update too large rather than causing a crash.
Derive the divisor instead of memorising it: state the target as the mean loss over the whole group, then show which factor each reduction convention needs. Be able to name the resulting effective-learning-rate change.
Demonstrate the verification habit — comparing a one-pass gradient against the accumulated one — and be alert to the variants that survive review, such as unequal micro-batch sizes and a short final group.
Frame it as an interface problem: the loss reduction convention is an implicit contract between data pipeline, loss and training loop, and it should be asserted in a test rather than trusted, because the failure is silent and mimics a bad learning rate.
## The invariant to hold on to There is one rule, and every convention question reduces to it: > The gradient you finally step with must be the gradient of the **average loss over all samples in the group**, the same quantity a single large batch would have produced. Write the group as `B = N * m`: `N` micro-batches of `m` samples. The target is `g = (1/B) * sum over all B samples of grad(loss_i)` Accumulation computes `g` as a sum of `N` partial contributions. The only question is what each contribution must be scaled by so the sum lands on `g`. ## The two conventions **Mean reduction (the common default).** Each micro-batch's loss is `(1/m) * sum of its m per-sample losses`. Its gradient is the mean gradient of those `m` samples. Adding `N` such means gives `N` times the group mean — a mean of means is the group mean only if you divide by the number of means. So you divide each micro-batch loss by `N` before its backward pass, and the accumulated buffer holds exactly `g`. **Sum reduction.** Each micro-batch's loss is the raw sum over its `m` samples. Accumulating all `N` of them gives the sum over all `B` samples, so you divide once by `B = N * m`. Dividing by `N` here — the reflex, because "N is the accumulation count" — leaves a gradient `m` times too large. The failure with real numbers: eight micro-batches of four samples, per-micro-batch losses summed, total divided by eight. You have divided by the number of micro-batches when you needed to divide by thirty-two, so every update is four times its intended size. Nothing errors. The effective learning rate is silently quadrupled. ## Why this bug is nasty A constant factor `c` on the gradient is not a subtle numerical issue; it is a change to the optimization recipe: - With plain stochastic gradient descent, gradient times `c` is exactly learning rate times `c`. A four-times step often diverges immediately, or trains but never reaches the loss the recipe promised. - With an adaptive optimizer such as Adam, part of the factor cancels, because the update divides by the square root of a second-moment estimate of the same gradients. That makes the bug *harder to see* — early behaviour looks nearly normal — while decoupled weight decay, which is not scaled by the gradient magnitude, now has a different relative strength than intended. "It looked fine so the scaling must be fine" is not evidence. - Gradient-norm-based safeguards read the wrong magnitude, so any threshold tuned on the correct scale now fires far too often or never. ## Unequal micro-batches Dividing by `N` is only correct when the micro-batches contain the same number of samples. When the last micro-batch of a dataset is short, or when samples are variable-length and micro-batches hold different token counts, mean-of-means silently up-weights the samples in the smaller micro-batch: each micro-batch contributes `1/N` of the update no matter how many samples it represents. The robust formulation is to weight each micro-batch by its share of the group: contribute `(m_j / total) * mean_j`, which is the same as summing per-sample losses across the whole group and dividing once by the true total count. For a token-level objective, `total` is the number of counted tokens, not the number of sequences. The related end-of-epoch trap is a short final group: seven micro-batches arrive instead of eight, the code divides by the hard-coded eight anyway, and that last update lands at seven-eighths of its intended magnitude. It is small and it is real; the fix is to divide by what the group actually contained, or to drop the ragged remainder. ## How to verify it rather than reason about it The check that settles the argument takes minutes: take a batch small enough to fit whole, compute its gradient in one pass, then compute the same batch as `N` micro-batches through your accumulation path, and compare the resulting gradient tensors elementwise. They should agree to floating-point tolerance. If they differ by a clean factor of `N` or `m`, you have found the divisor bug. If they differ in a way no single factor explains, look for something batch-dependent inside the forward pass instead. A cheaper smoke test: the loss value you log per update should be on the same scale as the loss of a single non-accumulated batch. If your logged loss jumps by a factor of `N` the day you turn accumulation on, the scaling is wrong — or, at minimum, your logging is accumulating the unscaled quantity.
- Your micro-batches hold different numbers of samples. What is the correct weighting?Weight each micro-batch by its share of the group rather than giving each an equal `1/N`. Equivalently, sum per-sample losses across the whole group and divide once by the true total count. With equal-sized micro-batches the two agree; with a short or variable-length one, equal weighting quietly gives the smaller micro-batch's samples more influence per sample.
- How would you prove your accumulation path is scaled correctly?Take a batch that fits in one pass, compute its gradient directly, then run the same samples through the accumulation path as several micro-batches and compare the two gradients elementwise. They should match to floating-point tolerance. A clean factor of the micro-batch count or the micro-batch size in the difference identifies exactly which divisor is wrong.
- Why is this bug easier to miss with an adaptive optimizer than with plain SGD?An adaptive method divides the update by a running estimate of the gradient's own magnitude, so a constant factor on the gradient largely cancels in the step. Training does not blow up the way it would under SGD, but anything not scaled by gradient magnitude — decoupled weight decay, norm thresholds, logged gradient norms — is now calibrated against the wrong scale.
saying these in an interview costs you the question
- Divides by the micro-batch count regardless of the loss reduction
- Assumes a mean of means is the overall mean
- Thinks a wrong factor would raise an error
- Ignores that the last group may be short
- Says adaptive optimizers make the scaling irrelevant