Accumulation changes when the optimizer steps; a correct large-batch approximation also preserves the weight of every example, including the final short microbatch.
Gradient accumulation and effective batch accounting
Define the unit being averaged
A microbatch is one forward and backward pass. An optimizer update can combine several microbatches; its effective batch is the total examples whose gradients contributed to that update, multiplied by the number of workers when gradients are reduced across workers. Dividing every microbatch mean by the number of microbatches fails when their sizes differ. For sample-averaged cross-entropy, sum the losses in each microbatch and divide by the total accepted examples in the update. Sampling policy determines which examples enter that denominator.
Step once at the accumulation boundary
Clear gradients before the first microbatch, call backward once per microbatch, then clip if required and step once. Clearing inside the loop discards earlier contributions. Stepping after each backward changes both the effective batch and the parameter point at which later gradients are computed. A final partial accumulation window still needs an update, with its own denominator; dropping it changes the training population. Count optimizer updates separately from examples seen and loader iterations.
State where equivalence ends
With a deterministic model and a loss separable over examples, accumulated gradients can match one full-batch gradient up to floating-point order. Batch normalization reads microbatch statistics, so it can change the function itself. Dropout draws different masks, and distributed reduction conventions can add another scaling factor. Learning-rate schedules indexed by optimizer updates should advance once per effective batch, not once per microbatch. Validation must compare the resulting model, not merely the first loss.
Control long-sequence cost
Accumulation reduces the peak activations of one pass but does not reduce the model weights, optimizer moments or the largest single-example activation. If one sequence alone exceeds memory, accumulation cannot fix it; trim sequence length, checkpoint activations or change the model. Do not retain each loss tensor in a Python list while logging; detach its scalar value so previous computation graphs can be released.
Prove the accounting
Run a fixed three-example batch as one pass, then as microbatches of two and one with copied initial weights. Compare every parameter gradient before stepping. Test an incomplete final window and a distributed configuration with a documented loss reduction convention. Record example count, microbatch count, optimizer-step count and gradient norm. A run that says only "batch size 47" conceals the update cadence that determines its training behavior.
Implementation
import copy
import torch
from torch import nn
torch.manual_seed(47)
full_model = nn.Linear(3, 2)
split_model = copy.deepcopy(full_model)
receipt_features = torch.tensor([[0.7, 0.2, 0.1],
[0.1, 0.6, 0.3],
[0.4, 0.1, 0.5]])
quality_labels = torch.tensor([0, 1, 0])
full_loss = nn.functional.cross_entropy(
full_model(receipt_features), quality_labels, reduction="mean")
full_loss.backward()
accepted_examples = quality_labels.numel()
for first, last in ((0, 2), (2, 3)):
logits = split_model(receipt_features[first:last])
loss_sum = nn.functional.cross_entropy(
logits, quality_labels[first:last], reduction="sum")
(loss_sum / accepted_examples).backward()
for full_parameter, split_parameter in zip(full_model.parameters(),
split_model.parameters()):
torch.testing.assert_close(full_parameter.grad, split_parameter.grad,
rtol=1e-6, atol=1e-7)Performance and operating cost
Each microbatch stores activations for at most its own samples, so peak activation memory is closer to the microbatch footprint than the full effective-batch footprint. Total forward and backward work remains roughly proportional to all examples; more microbatches add launch and synchronization overhead. The model and optimizer state remain resident. A short last batch must be weighted by its example count.
Common Mistakes
- Do not divide each unequal-size microbatch mean by the number of microbatches.
- Do not reset gradients or advance the scheduler between microbatches of one update.
- Do not promise full-batch equivalence when batch-dependent layers change their statistics.
Read next
- Batching and class sampling: know the population the optimizer sees
- Logits, cross-entropy and gradients: align the training calculation
- Mixed precision, loss scaling and gradient clipping
- Checkpoint recovery: save optimizer state and the run boundary
- Project: train a receipt model under a memory budget
Continue the workflow: Residual blocks and normalization state.
Continue the workflow: Recurrent hidden-state resets and truncated gradients.
Continue the workflow: Distributed sampler tails and global loss weighting.
Continue the workflow: AdamW decay groups and update-indexed schedules.
Continue the workflow: Project: measure activation recomputation on receipt training.
