Home›Journal›This post

Gradient Accumulation: Match the True Batch

Conserve token-loss and gradient numerators across variable-length microbatches, then align synchronization, scaling, clipping, optimizer, and scheduler clocks.

JP
JP Casabianca
AI Engineer and Product Designer · full-stack delivery · Bogotá

Gradient accumulation is only equivalent to a larger batch when every microbatch contributes to the same numerator and denominator. This guide builds a token-level accumulation receipt, catches mean-of-means drift, and aligns DDP, AMP, optimizer, and scheduler clocks before a training run is trusted.

Gradient accumulation starts with one invariant

A full batch defines the reference arithmetic: add every valid token loss, add every token’s gradient contribution, and divide each total by one count of valid prediction targets. Gradient accumulation should preserve that ledger while splitting memory work across microbatches. The number of loop iterations is operational detail, not a replacement denominator.

This matters whenever sequence lengths or masks differ. Dividing each microbatch by its own token count and then averaging those values gives every microbatch equal weight, even if one contains six useful targets and another contains one. The larger batch would weight the first six times more. A correct implementation preserves the shared numerator until the update boundary, where a single denominator represents the effective batch size in prediction targets.

Arithmetic equivalence is narrower than bitwise identity. Floating-point reduction order, dropout state, normalization layers, precision, and optimizer details can still differ. The goal is an inspectable necessary condition: the full-batch reference and the accumulated path use the same normalized records and agree within declared absolute and relative tolerances. Freeze those comparison conditions in the receipt: software revisions, reduction convention, random-state policy, precision, ignore index, and tolerance. A later run can then distinguish a changed objective from an expected numerical perturbation.

The QLoRA memory receipts help decide whether an adapter run fits. This receipt asks the next question: does its gradient accumulation still describe the same objective after the batch is partitioned? Treat the answer as a conservation check before interpreting loss curves, speed, or model quality.

Token numerator and denominator ledgerThree unequal microbatches conserve loss sums over ten targets to reach 0.64, while a crossed mean-of-means branch reaches 0.533333.CONSERVE THE NUMERATOR · COUNT VALID TARGETSMICRO 13 targets · Σℓ 1.2MICRO 26 targets · Σℓ 4.8MICRO 31 targetΣℓ 0.4(1.2 + 4.8 + 0.4) ÷ (3 + 6 + 1)= 0.64 ✓mean(0.4, 0.8, 0.4)0.533333 ✕
Token numerator and denominator ledger. The diagram and visible semantic equivalent carry the same conclusion.
Frozen 3/6/1 loss fixture
MicrobatchValid targetsLoss sumLocal mean
131.20.4
264.80.8
310.40.4

Conserved result: loss sum 6.4 ÷ 10 valid targets = 0.64. Rejected result: mean of local means = 0.533333….

Reading rule: labels, shapes, patterns, markers, and the semantic content carry every conclusion; color is supplementary.

Count prediction targets, not padded slots

For causal language modeling, labels are shifted: a sequence of length L usually contributes at most L - 1 next-token predictions. Attention padding, prompt-only masks, ignored labels, packed-sequence boundaries, and task-specific masking remove more. The denominator must therefore come from labels that actually participate in the loss, not tensor area, examples, or an advertised maximum sequence length.

Compute each microbatch loss with sum reduction, or recover an exact summed numerator from a documented reduction. Record the count of labels unequal to the ignore index beside it. Accumulate loss numerators and gradient numerators through the entire update window; divide by the window target count once. That is the token-normalized loss contract.

The official Transformers gradient-accumulation guide explains the variable-token denominator issue and the library mechanisms that address it. Pin the relevant library version because model and trainer hooks evolve. A receipt should name the formula it checked rather than infer behavior from a configuration field such as gradient_accumulation_steps.

Sequence packing without leakage adds another boundary: packed examples can share storage without sharing attention or targets across documents. Count after every mask and shift has been applied. If a window has no valid targets, reject it before backward; silently substituting one for a zero denominator turns invalid data into a plausible number.

Reject the mean of microbatch means

Work the smallest counterexample by hand. Three microbatches contain [3, 6, 1] valid targets and summed token losses [1.2, 4.8, 0.4]. The conserved numerator is 6.4; the denominator is 10; the window loss is 0.64. Every token contributes one share of that quotient.

The tempting alternative computes 1.2/3 = 0.4, 4.8/6 = 0.8, and 0.4/1 = 0.4, then averages those three values to obtain 0.533333…. It answers “what was the average microbatch mean?” rather than “what was the average valid-token loss?” The absolute drift is about 0.106667, entirely created by the denominator choice.

The same defect reaches gradients. Scalar gradient numerators [1.5, -0.6, 0.1] sum to 1.0, so the correct normalized result is 0.1. Individually normalizing and averaging gives (0.5 - 0.1 + 0.1) / 3 = 0.166667…. A loss log can therefore look smooth while the parameter update follows a different weighting.

Equal-size microbatches are an important control: when every valid-target count is identical, both formulas agree. Include that case so a test suite distinguishes the general rule from the fixture. Gradient accumulation needs both a failing unequal-size example and a passing equal-size control.

Put every clock on the update boundary

One window in the fixture uses K=3. It has three forward/backward calls but exactly one optimizer step, one scheduler step, one scaler update, and one post-step zeroing event. These counts define the update clock. A logging counter may tick per microbatch, but it must not masquerade as the optimizer’s global step.

Zero gradients before the first window or immediately after a completed step, never between numerator contributions. If a nonfinite check skips an optimizer step, document whether the scheduler and global update counter also remain still. Most training recipes intend those clocks to describe successful parameter updates, yet framework integrations can differ. The receipt makes the choice visible.

Learning-rate schedules are especially easy to advance too quickly. A scheduler configured for optimizer updates but called after every backward compresses its intended course by K. Checkpoints then serialize misleading step numbers, and resume logic can re-enter at the wrong learning rate. Align checkpoint naming, evaluation cadence, and scheduler state with the same declared boundary.

Activation recomputation changes compute and memory without changing this ledger. The activation checkpointing trade-offs remain separate from target normalization. A trustworthy gradient accumulation trace records both clocks: backward index for operational diagnosis and optimizer-update index for model-state history.

Five clocks on one update boundaryFive labeled rows show backward on all three microbatches but DDP synchronization, AMP unscale and clip, optimizer and scaler, scheduler and zero only at the third.K = 3 · ONLY THE LAST BACKWARD CLOSES THE WINDOWBACKWARDRUNRUNRUNDDP SYNC——1×AMP UNSCALE / CLIP——1×OPTIMIZER / SCALER——1×SCHEDULER / ZERO——1×MICRO 1MICRO 2MICRO 3 · STEP
Five clocks on one update boundary. The diagram and visible semantic equivalent carry the same conclusion.
  1. Backward runs three times, once per microbatch.
  2. no_sync covers forward and backward for microbatches 1–2; DDP sync occurs on backward 3.
  3. AMP unscale and optional clipping occur once after the final synchronized backward.
  4. Optimizer step and scaler update occur once at the effective-batch boundary.
  5. Scheduler advances once and gradients are zeroed after the step.

Reading rule: labels, shapes, patterns, markers, and the semantic content carry every conclusion; color is supplementary.

Place DDP synchronization deliberately

DistributedDataParallel normally synchronizes gradients during backward. For non-final microbatches, wrap both the forward and backward work in the model’s no-sync context so those local contributions can accumulate without a collective. The last backward in the window must run outside that context and trigger synchronization. Synchronizing every microbatch changes communication cost; synchronizing none leaves ranks with divergent updates.

The PyTorch DistributedDataParallel documentation explicitly warns that the forward pass must be inside no_sync() for the optimization to take effect. Record DDP no_sync use as a sequence of booleans, not as a claim that the run was faster.

Global normalization also needs global counts. Each rank can accumulate local loss and gradient numerators, but the denominator must represent the intended global window. Account for DDP’s gradient reduction convention when scaling; a sum and an average across ranks are not interchangeable. Test a two-rank fixture with matched window membership, then reject one whose rank-local tail decisions differ.

Distributed checkpoint recovery must preserve the next microbatch, partial numerator state if carrying a tail, scaler state, and optimizer clock. The distributed checkpoint recovery checklist helps define that resume boundary. Gradient accumulation is correct only when all ranks agree which backward is final and which denominator belongs to it.

Unscale and clip once per effective batch

Mixed precision introduces another clock but not another objective. Keep the scale factor consistent across every backward contribution in one effective batch. After the final synchronized backward, unscale gradients once, inspect them for nonfinite values, clip once if the recipe requires clipping, then call the optimizer and scaler update at the update boundary.

The official PyTorch AMP examples describe this mixed precision accumulation sequence. Unscaling early prevents later scaled gradients from being safely added to the same storage. Clipping each microbatch separately also changes the vector: the sum of clipped pieces is generally not the clipped sum.

Record the clipping norm, threshold, unscale count, nonfinite decision, step outcome, and scaler-update count. Do not turn those fields into performance or convergence claims. A receipt can establish that operations happened in a consistent order; it cannot establish that the chosen threshold is optimal or that lower precision matches full precision.

The full-batch comparison should use tolerances. Absolute tolerance protects values near zero; relative tolerance scales with magnitude. Report both deltas and the exact acceptance expression. Reduction order, fused kernels, and precision can create small differences even when gradient accumulation represents the same weighted objective. “Within this declared tolerance under this software stack” is an honest result; “identical training” is not.

Short final-window policy treeSeven microbatches group into three, three, and one before branching to flush, carry, or drop with distinct denominator and resume obligations.7 MICROBATCHES · K = 3 · DECLARE THE TAIL1234567window 1 · 3window 2 · 3tail · 1FLUSHuse actual targetsCARRYserialize partial stateDROPrecord discarded IDsALL RANKS MUST CHOOSE THE SAME DISPOSITION
Short final-window policy tree. The diagram and visible semantic equivalent carry the same conclusion.
Tail-policy obligations
PolicyAllowed useData consequenceRequired receipt field
FlushProcess the short final windowOne update using its actual valid-target countTail window membership and denominator
CarryContinue into the next segmentPreserves pending contributionsSerialized numerator, targets, IDs, and resume state
DropDeliberately omit the tailDiscards named microbatches and targetsDiscarded IDs/count and shared rank decision

Reading rule: labels, shapes, patterns, markers, and the semantic content carry every conclusion; color is supplementary.

Define the short final window

Seven microbatches with K=3 form two complete windows plus one tail. A flush policy processes [3,3,1]: the final update uses the last microbatch’s actual valid-target denominator and still produces one synchronized optimizer boundary. It must never be divided as though three microbatches existed.

A carry policy preserves the partial window for the next epoch or stream segment. Its checkpoint must include window membership, accumulated numerators, target count, backward count, and random-state boundary; otherwise resume silently changes the batch. A drop policy discards the tail deliberately and records the affected microbatch IDs and target count. Neither policy is universally best.

Every rank must choose the same disposition. If one rank flushes while another carries, their collective schedule can diverge or hang. Validate complete window manifests before allocating result arrays or entering a collective. The local lab rejects missing windows, duplicate IDs, mismatched gradient dimensions, unsafe counts, nonfinite numerators, inconsistent K, and rank-local tail decisions.

Gradient accumulation receipts should state flush, carry, or drop as data, not leave the behavior implicit in a loop ending. Short datasets, filtered examples, and failure recovery all exercise this boundary. A tail test belongs beside the happy path, not in a comment that production never reaches.

Publish an accumulation receipt

A useful receipt begins with schema and formula versions, then lists normalized ranks, microbatch IDs, window membership, valid-target counts, loss numerators, and gradient numerator vectors. It includes the global sums, normalized full-batch and accumulated views, naïve mean-of-means counterexample, absolute and relative deltas, tolerance, tail disposition, and every update-clock count.

Hash canonical JSON with sorted object keys and finite normalized numbers. Keep model files, prompts, labels, and training examples out of the artifact; they are unnecessary for this arithmetic and create privacy risk. The input hash binds the synthetic or redacted manifest, while a separate receipt hash binds the result and its formula version.

Eligibility is not success. Reject malformed data before computation, label arithmetic agreement clearly, and carry limitations beside the result: no autograd execution, no distributed collective, no AMP kernel, no optimizer update, and no evidence about throughput or convergence. Gradient accumulation is a conservation ledger, and the receipt proves only the ledger it actually evaluated. Archive rejected manifests with synthetic identifiers and their validation reasons so recurring contract mistakes become visible without retaining training content.

Use the frozen 3/6/1 fixture as a regression test, the equal-size control as a sanity check, and one real redacted manifest as the operational bridge. When all three agree, the denominator, communication boundary, and optimizer clock become inspectable instead of assumed. Keep the failing counterexample in the suite even after the implementation passes: it explains exactly what the guard prevents and stops a future refactor from replacing token weighting with a cosmetically reasonable batch mean.

Runnable local artifact — The lab does not execute autograd, collectives, AMP, an optimizer, or a training run; arithmetic agreement is necessary but not sufficient for numerical or convergence equivalence.

Plain text1 line
Validate every rank and microbatch before allocation, conserve global numerators and valid targets, declare the tail policy, and export the versioned receipt.