Does data-parallel replication let you train a model that does not fit on one device?
answer
- Count what each device must hold
- Replicated versus per-sample memory
- Only activations follow the micro-batch
- Weights, gradients, moments - all duplicated
- Throughput scales, capacity does not
basics
~20 sNo. Data parallelism puts a full copy of the weights, gradients and optimizer state on every device, so the model's own memory cost is unchanged. Only the activations shrink, because each replica sees a smaller micro-batch.
solid answer
~50 sNo - replication is the opposite of relief. Per-device memory has four consumers: parameters, gradients, optimizer state, and the activations stashed for the backward pass. Data parallelism duplicates the first three on every device unchanged; only the fourth shrinks, roughly in proportion to the per-device micro-batch. So going from one device to eight buys throughput, not capacity. For a model with P parameters trained in 32-bit with a two-moment optimizer, the fixed part is about 4P bytes of weights plus 4P of gradients plus 8P of moments - around 16 bytes per parameter, or roughly 21 GB for 1.3B parameters, before a single activation. Every device pays that in full, and the moment buffers are bit-for-bit the same numbers on all of them - which is exactly the waste that motivates strategies that stop replicating that state.
go deeper
Remember the direction: replication means every device carries the entire model, so more devices does not mean more room. What shrinks is only the number of samples each device processes at once.
Be able to itemize per-device memory - parameters, gradients, optimizer moments, activations - and say which of them depends on the micro-batch. The rough 16 bytes per parameter for a 32-bit two-moment setup is worth being able to derive.
Diagnose the failure before choosing a tool: test whether a micro-batch of one fits. Know that shrinking the micro-batch to buy memory raises the communication share of each step, and be able to price that trade.
Own the argument that identical optimizer state on every replica is redundancy you are paying for in hardware. Be able to frame when a team should move beyond plain replication and what that decision costs in complexity and debuggability.
## The four consumers of device memory During training, device memory goes to: 1. **Parameters** - the weights themselves. 2. **Gradients** - one buffer the same shape as the parameters. 3. **Optimizer state** - for a plain momentum optimizer, one extra buffer the size of the parameters; for an Adam-family optimizer, two (a first and a second moment). 4. **Activations** - the intermediate tensors saved during the forward pass because the backward pass needs them, plus transient workspace. Data parallelism replicates the model, so consumers 1-3 appear **in full on every device**. Only consumer 4 depends on how many samples that device is processing, and only consumer 4 shrinks when you split the batch. ## The arithmetic Take `P` parameters trained in 32-bit with a two-moment optimizer: - weights: `4P` bytes - gradients: `4P` bytes - two moment buffers: `8P` bytes That is `16P` bytes of fixed cost, per device, regardless of how many devices you have. For `P = 1.3e9` that is about 21 GB before any activation is stored. Adding a second, eighth, or sixteenth replica changes none of it - each new device must independently hold the same 21 GB. Activation memory, by contrast, scales roughly linearly with the per-device micro-batch (and with depth, and with resolution or sequence length). Moving from one device at batch 256 to eight devices at batch 32 divides the activation footprint by eight on each device. If activations were what overflowed, data parallelism at a smaller per-device batch genuinely helps. If the *model* was what overflowed, it does not help at all. ## Why the distinction gets confused The confusion is understandable: the practical symptom of both problems is the same out-of-memory failure. The diagnostic question is whether the run fails at a micro-batch of one. If a single sample already overflows, no amount of batch splitting will save you, because a micro-batch of one is the floor of what data parallelism can do. If it survives at micro-batch one but not at 32, the pressure is activations and the batch split is a legitimate lever - as is recomputing activations in the backward pass instead of storing them, which trades compute for memory without touching how many devices you use. ## The duplication argument Notice what the replicated optimizer state actually is. After the all-reduce, every device holds the same averaged gradient and applies the same update to the same weights, so the moment buffers on all `N` devices contain **bit-identical numbers**. For a 1.3B-parameter model that is roughly 10 GB of moment buffers per device holding exactly the same values as the 10 GB on every other device. It is pure duplication: none of the redundancy contributes anything computationally, it exists only because each device was given a whole model. That observation is the standard motivation for going beyond plain replication - if the state is identical everywhere, each device could hold only its slice of it and exchange what it needs. Where and how that slicing is done is a separate design question with its own tradeoffs; for this topic the point is simply that plain data parallelism does not do it, and that its memory cost per device is flat in the number of devices. ## What to say in an interview The crisp version is: data parallelism scales **throughput**, not **model size**. The samples-per-second curve improves with more devices (up to communication limits); the bytes-per-device curve does not move at all for parameters, gradients and optimizer state. If the interviewer's scenario is a model that will not fit, replication is the wrong tool and saying so quickly is the answer they are testing for. ## A second-order effect worth mentioning Smaller per-device micro-batches do not only reduce activation memory - they also reduce the compute per step on each device while leaving the per-step gradient communication unchanged. So the memory relief you buy by shrinking the micro-batch is paid for in scaling efficiency: the reduction becomes a larger fraction of each step. That coupling - memory down, communication fraction up - is the tradeoff a senior candidate is expected to notice rather than treat the micro-batch as a free knob.
- Going from one device at batch 256 to eight devices at batch 32, what actually shrinks?Only the activations stored for the backward pass and the transient workspace, both roughly proportional to the per-device micro-batch - so about eight times smaller on each device. Parameters, gradient buffers and optimizer moments are unchanged, because every replica still holds the whole model. If your out-of-memory failure came from the model rather than the batch, nothing improves.
- What if a single sample's activations already overflow the device?Then data parallelism has nothing left to give - a micro-batch of one is its floor. At that point you either trade compute for memory by recomputing activations in the backward pass instead of storing them, reduce resolution or sequence length, or stop replicating and split the model itself across devices. The choice is about which resource you have spare.
- Why is replicated optimizer state described as pure waste at scale?Because every replica applies the same averaged gradient to the same weights, the moment buffers end each step containing identical numbers on every device. For a billion-parameter model that is several gigabytes per device duplicating what its neighbours already hold, contributing nothing computationally. It exists only as a consequence of handing each device a whole model.
saying these in an interview costs you the question
- Says adding replicas lets a larger model fit
- Forgets optimizer state when accounting for memory
- Thinks gradients are freed before the optimizer step
- Assumes memory per device falls as devices are added
- Cannot distinguish an activation overflow from a parameter overflow