Why does halving the training batch size cut device memory when the weights are unchanged?
answer
- Two piles, only one is per-sample
- Parameter count fixes the first pile
- Gradients are one per weight, not per sample
- Activations carry a batch axis
- Saving depends on the ratio between piles
basics
~20 sWeights, gradients and optimizer state are sized by the parameter count, not by batch size. Activations, the per-layer outputs held for the backward pass, exist once per sample, so halving the batch roughly halves that share.
solid answer
~50 sTraining memory splits into two piles. The first is **model state**: one weight tensor per parameter, one gradient tensor per parameter, and whatever the optimizer keeps per parameter — a momentum buffer, or Adam's two moment estimates. All of that is sized by the parameter count and is completely indifferent to batch size. The second pile is **activations**: the outputs each layer produces on the forward pass and holds until the backward pass consumes them. Those exist once per sample, so they scale roughly linearly with the number of samples in the batch. Halving the batch halves the second pile and leaves the first untouched. Whether that is a big win depends on the ratio: for a wide vision model at a healthy batch size activations dominate and the saving is large; for a small-input model whose optimizer state is most of the footprint, halving the batch barely moves the number.
go deeper
Be ready to name the four things in memory during a step: weights, gradients, optimizer buffers and activations, and to say which one moves with batch size. That single split is the whole answer at this level.
Explain why gradients are one per parameter rather than one per sample, and estimate the split from a measured total: given 20 GB at batch 32 with 12 GB of state, say what batch 16 costs.
Show that you check the ratio before recommending a batch cut. An interviewer expects you to say when batch size is the wrong lever — a parameter-heavy model where activations are a small slice — and to name the batch-size-one floor.
Own the framing that memory budget is an architecture decision, not a runtime knob. Be ready to argue when a team should re-shape inputs or width instead of shrinking batches into a regime where throughput and batch statistics both degrade.
## The two piles Everything that occupies device memory during a training step falls into one of two categories, and the whole answer follows from knowing which category each thing is in. **Pile one — model state.** This is memory whose size is a function of the parameter count and nothing else: - the **weights** themselves, one number per parameter; - the **gradients**, one number per parameter, because the backward pass produces exactly one partial derivative per weight; - the **optimizer state**, which is whatever the update rule carries between steps. Plain stochastic gradient descent carries nothing. SGD with momentum carries one velocity buffer per parameter. Adam and AdamW carry two: a first moment estimate and a second moment estimate. The critical point for this question is that **none of these depends on how many samples you push through**. A gradient is not accumulated per sample and stored per sample; the per-sample contributions are summed or averaged into a single gradient tensor as the backward pass runs. Sixteen samples and one sample produce a gradient tensor of exactly the same shape. **Pile two — activations.** On the forward pass, each layer computes an output tensor from its input. Most of those intermediate tensors have to be kept alive until the backward pass reaches that layer, because computing the gradient with respect to a layer's weights needs the values that flowed into or out of it. Those retained tensors are the activations. An activation tensor's shape starts with the batch dimension. A convolutional feature map is (batch, channels, height, width); a sequence model's hidden states are (batch, length, width). Every retained tensor carries that leading batch axis, so the sum over all of them is, to a very good approximation, `per_sample_bytes * batch_size`. Halve the batch and you halve it. ## Doing the arithmetic Suppose a run at batch 32 occupies 20 GB, and you know the model state is 12 GB. Then activations are about 8 GB, or 250 MB per sample. At batch 16 you expect roughly `12 + 4 = 16 GB` — a 20% saving, not a 50% saving. The mistake juniors make is assuming total memory is proportional to batch size; only one of the two piles is. The same arithmetic run the other way tells you the ceiling. As batch size goes to one, memory does not go to zero: it converges to `model_state + one_sample_activations`, plus small fixed overheads for temporary buffers a layer needs while computing. If that floor already exceeds the device, no amount of batch reduction rescues the run — the shape of a single sample has to change. ## Where the ratio comes from Which pile dominates depends on the architecture, and the two extremes look completely different: - A **large language-style stack or a very wide network trained on small inputs** is parameter-heavy. Model state can be tens of gigabytes while activations at a modest batch are a few. Here batch size is a weak lever. - A **convolutional network on high-resolution images or volumes** is activation-heavy. The parameter count may be tens of millions — under a gigabyte of state — while a single sample's retained feature maps run to several gigabytes because the early layers hold full-resolution tensors. Here batch size is the dominant lever, and the batch that fits may be a single-digit number. So "how much does batch size buy me" is not a fact about training in general; it is a fact about the ratio in your specific model, and a candidate who can name the ratio rather than a rule of thumb is answering at the right level. ## Secondary effects worth knowing A few details soften the clean linear picture. Some memory is allocated once and reused for scratch space regardless of batch size, so the activation share is linear plus a constant. Very small batches can also make the arithmetic on the device less efficient, so you pay in throughput for the memory you free. And normalization layers that compute statistics across the batch behave differently at small batch sizes — that is a statistical concern, not a memory one, but it is why shrinking the batch is not a free choice. The takeaway to carry into an interview: **parameter count sets a fixed floor; input shape and batch size set everything above it.** Being able to say which of your two numbers is which, and estimate both, is what the question is really testing.
- Does halving the batch halve total training memory?No. Only the activation share scales with batch size; weights, gradients and optimizer state are unchanged. If a 20 GB run has 12 GB of model state, batch 16 instead of 32 gets you to roughly 16 GB, not 10 GB. Estimate the ratio before promising anyone a number.
- At batch size one, what still sets a floor on memory?The full model state — weights, gradients and optimizer buffers — plus one sample's retained activations, plus small scratch buffers. That floor is a property of the parameter count and the input shape. If it exceeds the device, the shape of the model or the input has to change; there is no batch left to cut.
- Do gradients get bigger when the batch is bigger?No. The backward pass produces one partial derivative per parameter, and the per-sample contributions are summed into that single tensor as it goes. The gradient buffer has the same shape as the weights at any batch size. Only the intermediate activations feeding that computation are per-sample.
Model state is the workshop's fixed machinery; activations are the parts sitting on the bench mid-assembly. Ordering half as many parts clears the bench, but the machines take up exactly as much floor space as before.
saying these in an interview costs you the question
- Says training memory is basically the model size
- Thinks the gradient tensor grows with batch size
- Assumes total memory is exactly proportional to batch size
- Believes batch size changes the parameter count
- Cannot name what the optimizer stores per parameter