What fills a GPU’s memory during training?

Every GPU memory question I have been asked in an interview reduces to one calculation, and most candidates cannot do it on the spot. “Can you fine-tune a 70B model on eight H100s?” “How many sequences can one GPU serve at 32k context?” “Why does the job OOM in the backward pass and not the forward?”

The calculation is an inventory. Training keeps six kinds of thing in HBM and inference keeps two, each with a formula that fits on a line. Once you can write the inventory, every one of those questions is arithmetic, and every sharding technique in the field is a line item you are dividing by something.

The numbers below are for Llama 3 70B unless stated: 70.5 billion parameters, 80 layers, hidden size 8192, 64 attention heads, 8 key-value heads.

They will ask What occupies GPU memory when training a large model, and how do you fit a model that does not fit?
The 30-second version Four things. Model states: with Adam in mixed precision, sixteen bytes per parameter, made of a BF16 weight, a BF16 gradient, and FP32 master weights plus two Adam moments. Activations saved for the backward pass: per layer roughly 34 times sequence times batch times hidden size in bytes, plus a quadratic attention term that FlashAttention removes. Fixed overheads: the CUDA context, communication buffers, kernel workspace. And allocator fragmentation. The states shard for free across data parallelism with ZeRO: stage one divides the optimizer states by the GPU count, stage two the gradients too, stage three (FSDP) the weights as well, at the cost of gathering weights every step. Activations do not shard across data parallelism, so you cut them with tensor and sequence parallelism, pipeline parallelism, context parallelism, and activation recomputation. Inference drops the gradients and optimizer, and what grows instead is the KV cache: two times layers times KV heads times head dimension bytes per token, read in full on every decode step.

Why sixteen bytes per parameter?

Take one parameter and follow it through an Adam step in mixed precision.

The forward pass reads it as a BF16 weight, two bytes. The backward pass writes a BF16 gradient, two more. The update cannot be applied in BF16, because eight bits of mantissa would swallow a learning rate times a small gradient, so an FP32 master copy holds the true value, four bytes. Adam keeps two running moments per parameter, both FP32, eight bytes.

\[2 + 2 + 4 + 4 + 4 = 16 \text{ bytes per parameter}\]

Seventy billion parameters need 1.13 TB. Fourteen H100s, or six B200s, before the first token enters the model. Three quarters of it is optimizer state, which is why the two live ideas for shrinking it, 8-bit Adam and momentum-only optimizers like Muon, both attack that column.

Inference pays two of the sixteen. That single fact is why a model that serves on one GPU takes a rack to train.

How much do the activations cost?

The backward pass needs what the forward pass saw. Megatron’s accounting for one GPT-style layer in BF16 is:

\[\text{bytes per layer} = s\,b\,h\left(34 + 5\,\frac{a\,s}{h}\right)\]

for sequence length $s$, micro-batch $b$, hidden size $h$ and $a$ attention heads. The first term is every intermediate of the attention block, the MLP and the norms, linear in the sequence. The second is the attention matrix itself: the scores, the softmax output and the dropout mask, five bytes per score, quadratic in the sequence.

At 8k tokens and $h = 8192$ the linear term is 2.3 GB per layer and the quadratic term is 21 GB. Multiply by 80 layers and one sequence needs 1.9 TB of activations. That is why long-context training did not exist before FlashAttention, which never writes the attention matrix and stores one log-sum-exp per query row instead. With it, the same sequence needs 182 GB. Still more than a GPU, but now within reach of sharding.

Llama’s SwiGLU and grouped-query attention shift the constant 34 somewhat. The shape of the formula is what to remember: linear in tokens for everything, quadratic for attention, and the quadratic part is the part a better kernel deletes.

How do the states shard?

Plain data parallelism replicates the states on every GPU. The ZeRO paper’s observation was that the replicas are redundant: the optimizer only needs its own slice of the parameters to update them.

stage sharded bytes per GPU for 70B on 64 GPUs
DDP nothing $16\Psi$ 1.13 TB
ZeRO-1 optimizer states $4\Psi + 12\Psi / N$ 296 GB
ZeRO-2 + gradients $2\Psi + 14\Psi / N$ 156 GB
ZeRO-3, FSDP + weights $16\Psi / N$ 17.6 GB

Stage three gathers each layer’s weights just before the layer runs and frees them after, so the communication per step is about three times the weight bytes: two all-gathers, one for the forward and one for the backward, and a reduce-scatter for the gradients. For 70B that is 423 GB per GPU per step, and it has to hide behind compute, which is the overlap story from Part 1 again. PyTorch’s FSDP2 and Megatron’s newer FSDP implementation both exist to do that hiding well.

At sixty-four H100s, ZeRO-3 leaves 62 GB per GPU for activations. At eight, it leaves nothing. That is the answer to the fine-tuning question: you can, with the optimizer states offloaded or quantised, or with LoRA so there are almost no optimizer states at all, but not the full model with full Adam.

How do the activations shard?

They do not shard for free, because each data-parallel GPU keeps the activations of its own micro-batch. Four levers cut them.

Tensor parallelism splits each layer’s matmuls across $t$ GPUs, and with sequence parallelism, which splits the norms and dropouts along the sequence too, every activation is divided by $t$. Pipeline parallelism gives each GPU only its own layers, but keeps several micro-batches in flight to fill the pipeline, so it helps less than it looks. Context parallelism splits the sequence itself across GPUs, which is how 128k-token training fits at all.

The fourth lever is recomputation. Selective recompute drops the internals of the attention block and regenerates them in the backward pass, at a few percent extra compute. Full recompute stores only each layer’s input, $2sbh$, and recomputes everything else, at about a third more compute. Full recompute is what everyone did before FlashAttention and Megatron’s selective scheme made it mostly unnecessary.

Llama 3 405B trained with tensor parallelism of 8 inside a node, pipeline parallelism of 16 across nodes, context parallelism of 16 for the long-context stage, and data parallelism across the remaining GPUs with FSDP inside the data-parallel group. Every one of those numbers is a division of a line in the inventory.

Key idea States divide by the data-parallel count. Activations divide by tensor, sequence and context parallelism, and shrink with recompute. Every parallelism scheme is a choice of what to divide by what, and the communication it costs is the byte count you did not want to store.

What does the bill look like per GPU?

Add the shard of the states, the activations after sharding, and the overheads nobody puts in a formula. The CUDA context and libraries take around a gigabyte. NCCL holds communication buffers. cuBLAS and the attention kernels want workspace. And the caching allocator keeps blocks it cannot immediately reuse, ten to twenty percent on a busy training step.

Then the peak must fit, not the average. The backward pass of the last layer runs with every saved activation still resident, and that moment is when jobs die.

The figure’s fifth step does this sum for any combination. Two levers dominate it: how many GPUs share the states, and how long the sequences are. Everything else is fine-tuning, and the headroom left over is what lets you raise the micro-batch, which is what raises the arithmetic intensity from Part 2.

Fragmentation deserves one practical note. PyTorch’s allocator can be told to grow existing blocks instead of fragmenting new ones:

PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True

It is the first thing to try when a job OOMs with torch.cuda.memory_reserved() far above memory_allocated().

What fills memory when serving?

Inference drops the gradients and the optimizer. The weights cost two bytes per parameter in BF16, one in FP8, half in INT4. What grows instead is the KV cache.

Every token in every live sequence keeps a key and a value per layer:

\[\text{bytes per token} = 2 \cdot L \cdot n_{kv} \cdot d_h \cdot \text{bytes per value}\]

For Llama 3 70B in BF16 that is $2 \times 80 \times 8 \times 128 \times 2 = 320$ KiB. At 8k context a sequence carries 2.7 GB; at 128k, 43 GB, more than a whole H100. With all 64 heads as KV heads it would have been 2.5 MiB per token, and grouped-query attention is the single reason long-context serving is possible at all.

DeepSeek-V3 compresses further with multi-head latent attention: a 576-wide latent per token per layer instead of full keys and values, about 69 KiB per token across its 61 layers. FP8 caches halve it again. Paged attention, from vLLM, stores the cache in fixed blocks so that sequences of different lengths share HBM without fragmentation.

The cache is what caps how many sequences a GPU can hold, and, because every decode step reads all of it, it is the byte count that caps the arithmetic intensity from Part 2. Prefill fills it at compute-bound speed; decode reads it at memory-bound speed. The two phases want different hardware, and Part 6 has the systems that give it to them.

Which memory level holds which of these?

All of the above lives in HBM, because all of it must survive between kernels. That is the bottom of Part 3’s pyramid.

What the levels above hold is the working set of the kernel running right now. During a matmul, the current tiles of the weight and the activation are in shared memory and the accumulator is in registers. During attention, a block of queries and the passing blocks of keys and values are in shared memory, the block of scores is computed and consumed there, and the running maximum, running sum and output tile are in registers. During the optimizer step, a slice of the master weights and moments streams through registers once and goes back.

The distinction is lifetime. Anything that must be there when the next kernel starts is in HBM and in this inventory. Anything that only needs to exist during one kernel is above HBM, and the roofline never counts it. Part 5 is about the biggest thing that ever moved from the first category to the second.

They will ask The job fits in the forward pass and dies in the backward pass. Why, and what would you change first?
The 30-second version The forward pass accumulates saved activations layer by layer; the peak comes at the start of the backward pass, when every layer's activations are resident and the first gradient buffers are being allocated on top. So the forward fits and the backward does not. First check the sequence length and micro-batch, because activations scale with their product and nothing else does. Then enable activation recomputation, selective first since it is nearly free, or increase tensor or context parallelism to divide the activations. If the states rather than the activations are the problem, move up a ZeRO stage. And set the allocator to expandable segments before concluding anything, because fragmentation at the peak looks exactly like being out of memory.

Rapid fire: can you do these from memory?

  1. Derive sixteen bytes per parameter and say which four of them inference keeps.
  2. Write Megatron's activation formula and name what each term stores.
  3. Which term does FlashAttention remove, and what does it store instead?
  4. Give the bytes per GPU for ZeRO stages one, two and three, and the communication stage three adds.
  5. Why do activations not shard across data parallelism, and what four levers cut them?
  6. Where in the step is peak memory, and why?
  7. Compute the KV cache per token for Llama 3 70B in BF16 and the effect of GQA versus 64 KV heads.
  8. What is allocator fragmentation, how do you detect it, and what is the one-line mitigation?

Part 5 takes the largest object that ever left this inventory, the attention matrix, and follows the algorithm that keeps it above HBM: FlashAttention, from the first tiled version to the one written for Blackwell this year.