skip to content

Your weights fit on one device but Adam's moments do not — how do you shard the optimizer state?

level: seniorimportance: should knowfreq 40%

answer

  1. every replica computes the same update
  2. moments dominate the per-parameter bytes
  3. give each device a parameter slice
  4. reduce-scatter, update, then all-gather
  5. same wire volume as before

basics

~20 s

Partition parameters across replicas so each device keeps only its slice of Adam's moments and master weights, updates that slice, then all-gathers the updated weights. State memory falls by the replica count at no extra communication.

solid answer

~50 s

Adam keeps a first and a second moment per parameter, and mixed-precision training also keeps a full-precision master copy — roughly 12 bytes per parameter of state against 2 bytes for a half-precision weight. Every replica computes the identical update from the identical reduced gradient, so all of that state is replicated for no reason. Partition the parameters across the `N` replicas: each device owns `1/N` of the moments and master weights, receives only its slice of the reduced gradient, applies the update to its slice, and then the updated weights are all-gathered so everyone has the full model for the next forward pass. State memory drops by `N`. Communication is essentially free, because an all-reduce is already a reduce-scatter followed by an all-gather — you are just doing the update in between. Note this is not model parallelism: every device still runs the whole network.

go deeper

for a junior

Recall that Adam stores two extra numbers per parameter plus a full-precision copy, so optimizer state is several times the size of the weights and is usually what exhausts memory first.

for a middle

Explain why the state is redundant across replicas that all compute the same update, and describe the reduce-scatter, local update, all-gather sequence that removes the duplication.

for a senior

Demonstrate that you would reach for this before any model surgery because the wire volume is unchanged, and that you know the sharp edges: global-norm clipping, checkpoint format and resharding.

for a principal

Own the ordering of remedies under a memory constraint and the cost curve behind it — free state sharding first, then memory-for-bandwidth trades, then splitting the model, with the schedule and checkpoint implications of each made explicit.

## Why the optimizer state is the thing that overflows A parameter in a mixed-precision training loop is not one number in memory, it is several. There is the half-precision copy the forward and backward passes actually use, its gradient, and then the optimizer's bookkeeping: Adam maintains an exponential moving average of the gradient (the first moment) and of the squared gradient (the second moment), both per parameter, and mixed-precision training also keeps a full-precision master copy of the weight so that tiny updates are not lost to rounding. Those three full-precision buffers are the bulk. Roughly speaking the optimizer-related state is several times the size of the weight copy itself — which is why a model whose weights would comfortably fit can still fail to train on the same device. The second moment is what makes Adam Adam: the update divides the (bias-corrected) first moment by the square root of the (bias-corrected) second moment plus a small constant, giving each parameter its own effective step size. You cannot drop either buffer without changing the algorithm. ## The redundancy to exploit Run the same model on `N` devices with different data on each, and the gradients are combined so that every device ends the backward pass holding the *same* averaged gradient. Every device then feeds that same gradient into the same optimizer holding the same moments, and computes bit-for-bit the same update. `N` devices are doing one device's worth of work and storing `N` copies of one device's worth of state. That duplication is pure waste, and it is the thing to remove. ## The scheme Partition the parameter vector into `N` disjoint slices and assign slice `i` to device `i`. Then: 1. Each device runs the full forward and backward pass on its own data and produces a full gradient — nothing about the model changes. 2. Instead of an all-reduce, do a **reduce-scatter** of the gradients: device `i` ends up with the averaged gradient for *only* its own slice, and no one holds the whole averaged gradient. 3. Device `i` holds the first moment, second moment and master weights for its slice alone, and applies the optimizer update to that slice. 4. **All-gather** the updated half-precision weights so every device again has the complete model for the next forward pass. Optimizer-state memory per device falls by a factor of `N`, exactly. ## Why the communication is nearly free A ring all-reduce is not a primitive — it is implemented as a reduce-scatter followed by an all-gather, each moving roughly one parameter-vector's worth of data per device. The scheme above uses precisely those two collectives, with the optimizer step inserted between them. The volume on the wire is essentially the same as the replicated version; what changes is that the all-gather now carries updated weights rather than reduced gradients. This is the reason this technique is close to a free win, and it is why it should usually be the first thing you try when memory is short. ## What it is not This is **not** model parallelism. No weight matrix is split, no layer is assigned to a particular device, and every device still executes the entire network on every step. What is partitioned is the optimizer's private bookkeeping, plus, momentarily, the gradient. A full copy of the weights is still materialised on every device before the forward pass. So the technique relieves state pressure and nothing else — if a single layer's weights, or the activations one device must retain, still do not fit, sharding the moments will not save you and you have to split the model itself. The idea does extend further along the same axis. You can also keep the gradients partitioned rather than materialising a full gradient buffer, and, more aggressively, keep the *weights* partitioned too and gather each layer's parameters only just before that layer is used, discarding them immediately afterwards. That last step removes almost all of the per-device weight copy, but it adds a collective for every layer in both the forward and the backward pass — a real memory-for-bandwidth trade rather than a free one. ## Operational sharp edges **Gradient clipping.** Clipping by global norm needs the norm over *all* parameters, but no device holds them all. Each device computes the sum of squares over its own slice and those scalars are all-reduced before the square root is taken and the scale applied. Skip that and each device clips against its own partial norm, which silently changes the effective threshold and can quietly destabilise training. **Checkpointing.** The state now lives in `N` pieces. Either gather it on save to produce a device-count-independent checkpoint, or accept a sharded format and pay a resharding step whenever the device count changes. **Partition granularity.** Slicing at arbitrary byte offsets splits individual tensors across owners, which complicates any per-tensor logic; partitioning whole tensors is simpler but can leave the load uneven when one tensor dominates the parameter count. **Determinism.** The updated weights now arrive through a collective, so reduction order affects the last bits. That is normally harmless, but it makes bitwise run-to-run reproducibility harder to guarantee.

  • What breaks about gradient clipping once the optimizer state is sharded?
    Clipping by global norm needs a norm over all parameters, and no device holds them all. Each device must compute the sum of squares over its own slice and all-reduce those scalars before taking the square root and applying the scale factor. If you skip that, every device clips against its own partial norm, which silently raises the effective threshold and can destabilise training in a way that is very hard to spot.
  • When does sharding the state stop being enough, forcing you to split the model itself?
    When a full copy of the weights plus the activations a device must retain still does not fit, or when a single layer is too large for one device. This scheme leaves the whole model resident on every device during the forward and backward pass — it only removes duplicated bookkeeping. At that point you need tensor parallelism to split individual matrices, or pipeline parallelism to give each device only some layers.
  • Can you push the same idea to the gradients and the weights?
    Yes, along the same axis. Keeping gradients partitioned rather than materialising a full buffer removes another copy at no real cost. Partitioning the weights themselves is more aggressive: each device stores only a slice and must gather a layer's parameters just before using it and discard them after, in both passes. That nearly eliminates the resident weight copy but adds a collective per layer — memory traded for bandwidth, not a free win.

Eight accountants each keeping a complete, identical copy of the same ledger and making the same entries. Give each one an eighth of the pages, let them post only their own entries, and photocopy the updated pages to everyone once at the end of the day.

saying these in an interview costs you the question

  • Calls optimizer-state sharding a form of model parallelism
  • Claims it removes the weight copy from each device
  • Says it roughly doubles per-step communication volume
  • Forgets that global-norm clipping now needs its own reduction
  • Thinks Adam's second moment can simply be dropped to save memory

context