Activation Memory: Why the Forward Pass Costs More Than the Weights

Inside LLM Fine-Tuning, part 3 of 15: Activation Memory: Why the Forward Pass Costs More Than the Weights

Activation memory is the part of a training run’s memory bill that almost nobody counts and almost everybody trips over. Weights are easy to count, so weights get the attention. Activations are the ones that end the run with an out-of-memory error, OOM for short, and they do it for reasons that have nothing to do with how many parameters the model has.

An activation is an intermediate value: something the model computes on the way from input to answer, and then normally throws away. Training does not throw them away. It holds them, all of them, all at once, and that hoard is often larger than the model itself.

By the end of this part you will be able to work out a run’s activation memory by hand from four numbers, explain from first principles why the numbers have to be kept at all, and name every lever that reduces them. Everything is built up from scratch here. You do not need to know what a gradient, a tensor or a transformer block is before you start.

Tell an activation apart from a weight, using code you already write

Start somewhere completely familiar. Here is a function you could have written on any ordinary Tuesday.

def f(x):
    a = x * 3
    b = a + 2
    return b * b

Call it with x = 5. Inside, a becomes 15, b becomes 17, and the function returns 289. a and b are intermediate values. They exist only while the function is running. When it returns they are gone, and nobody misses them.

A neural network is that same idea, much longer. Data goes in at one end, a few hundred operations run in order, and an answer comes out the other end. Running the model in that direction, from input to answer, is called the forward pass. Every intermediate value the forward pass computes along the way is an activation.

That is the whole definition. An activation is an intermediate. It is not a setting you configured, and it is not something saved to disk.

The other kind of number in the building

In the function above, the 3 and the 2 are constants baked into the code. A neural network has the same category of number, except there are rather a lot of them and they were learned rather than typed. Those are the weights, also called parameters. Qwen2.5-1.5B, the model used throughout this series, has about 1.54 billion.

Weights and activations behave in opposite ways, and holding that difference clearly is most of the battle.

  • Weights. Their count is fixed the moment you choose a model. They are loaded once and stay alive for the whole run.
  • Activations. Their size is set by how much data you push through in one go. They appear and vanish inside a single step, and two of the numbers that set their size are training flags you type yourself.

Nearly every memory surprise in fine-tuning comes from applying weight intuition to activations. You can run out of memory on the exact same model by changing one number in a config file.

One more word: tensor

The intermediates in a neural network are not single numbers like a and b. They are tensors. A tensor is a block of numbers with a shape, and that is the entire concept.

A single number is a tensor with no shape. A list of 1,536 numbers is a tensor of shape 1536. A tensor of shape 4 x 1024 x 1536 is a three-level nested array: 4 blocks, each holding 1,024 rows, each row holding 1,536 numbers. If you have ever allocated a multi-dimensional array, you already know what a tensor is.

The one extra fact worth carrying is that a tensor is stored as one contiguous run of bytes, so its memory cost is exactly the number of elements multiplied by the bytes per element. No overhead, no surprises. We will use that constantly.

See why the backward pass cannot throw an activation away

If activations are only intermediates, why not free each one the moment the next operation has consumed it? That is exactly what happens when you run a model to get an answer. Training is different, and the reason is worth building carefully, because everything else here follows from it.

Training does three things in a loop. It shows the model an example. It measures how wrong the answer was. Then it nudges every weight slightly in whichever direction would have made the answer less wrong. Three pieces there, each with a name.

  1. The loss is a single number saying how wrong the answer was. Big is bad. All of training is an effort to make it smaller.
  2. The gradient of a weight is one number that answers one question: if I nudge this weight up by a hair, does the loss go up or down, and how fast? One gradient per weight, so about 1.54 billion of them for our model. If you want that idea developed slowly with pictures, the walkthrough of how a neural network learns builds it from nothing.
  3. The backward pass, also called backpropagation, is the procedure that computes all of those gradients. It starts at the loss and works back towards the input, which is where the name comes from.

Nothing there mentions memory yet. Here is where memory arrives.

The whole argument, in two numbers

Take the smallest imaginable step of a network. One number goes in, it is multiplied by one weight, one number comes out.

  • The input is x = 4.
  • The weight is w = 0.5.
  • So the output is y = w * x = 2.

The backward pass has already worked its way back through everything after this step, and it arrives here carrying exactly one number: how much the loss changes for each unit of change in y. Say that number is -1. Read it as “push y up by 1 and the loss comes down by 1″. So we would like y to be larger.

The question this step has to answer is: how should w change? Work it out by hand. Raise w from 0.5 to 0.6, a change of 0.1. Then y goes from 2 to 0.6 * 4 = 2.4, a change of 0.4. Raise it again from 0.6 to 0.7 and y becomes 2.8, another 0.4. Every single unit of change in w moves y by 4.

Four. Which is x.

The sensitivity of the output to the weight is the input. Combine that with the number arriving from above and you have this weight’s gradient: the loss changes by -1 * 4 = -4 for each unit of change in w. That is the instruction the optimizer acts on, the optimizer being the component that decides how far each weight actually moves once it knows which way to go.

Now the point of the whole exercise. That calculation needed x. Not the shape of x, not a summary of x, the actual value 4. Watch what happens if only the input differs. Keep w = 0.5, keep the same message from above, and set x = 100. The sensitivity is now 100, and the gradient is -100 instead of -4. Same weight, same message, an instruction twenty-five times larger, purely because the input was different.

There is no way to recover that number after the fact. Either the forward pass kept x alive in memory, or this weight’s gradient simply cannot be computed. That is the reason activations are stored. They are an operand of a multiplication that happens later.

The same idea with grids instead of single numbers

A real layer multiplies a tensor of inputs by a grid of weights, doing millions of multiply-and-add operations at once. The recipe does not change. The gradient for the weight grid is the message arriving from above, combined with the stored input tensor.

You will see it written compactly as dL/dW = dL/dy * x^T. Read it straight back in the words you just derived. dL/dW is the gradient for the weights. dL/dy is the gradient that arrived at this step’s output, the message from above. x is the stored input. The ^T means transpose, which is flipping a grid on its diagonal so its rows become its columns, bookkeeping to make the two grids line up for multiplication. It adds no new idea and no extra memory.

FORWARD → y = W · x stash x — backward will need it x held in VRAM… ← BACKWARD dL/dW = dL/dy · xᵀ → now x can be freed ✓

The forward pass stashes x because the backward pass needs it as the second half of a multiplication. Drop x and this weight’s gradient is not merely inaccurate, it is uncomputable.

If you want that derivation with every intermediate step spelled out, the full calculus behind backpropagation is there for you. It is not required reading, because the two-number version above is the same fact.

What actually gets saved, and why the byte count is approximate

The framework decides which intermediates to keep, and it is more selective than you might expect. As you run the forward pass, PyTorch, the library most training code is written in, quietly records each operation and the values that went into it. That recording is the autograd graph, and it behaves like an undo history. When you ask for gradients it replays the list in reverse, and each recorded operation asks for what it needs.

Values that some backward step will need are saved, meaning they stay resident in GPU memory instead of being freed. A layer that multiplies by a weight grid saves its input, for the reason you just worked through. Values nothing will need are dropped at once, and some cheap operations prefer to recompute their result later rather than store it.

So the exact byte count depends on the framework version, the attention implementation and your flags. What does not depend on any of that is the shape of the answer: roughly a dozen saved tensors per layer, each about the size of the data flowing through. Treat every activation figure in this part as order of magnitude. The arithmetic below is right to within maybe 30 percent, and 30 percent is enough to tell you whether a run fits.

Predict the exact moment in a step when memory peaks

Knowing that activations are kept is half the story. Knowing how long each one is kept is the other half, and it is what gives the memory curve its shape.

Take a network with three layers and follow the bookkeeping. During the forward pass, layer 1 runs and saves its input, call it x1. Layer 2 runs and saves x2. Layer 3 runs and saves x3. Nothing has been freed, because nothing has been used yet.

Then the backward pass starts, and it runs in reverse. It reaches layer 3 first, uses x3, and frees it. Then layer 2, uses x2, frees it. Then layer 1, uses x1, frees it.

Look at x1. It was created first and it is consumed last. In a 28-layer model, the first layer’s saved input has to survive 27 more forward layers and 27 backward layers before anything needs it.

forward → L1 → x1 L2 → x2 L3 → x3 stored: x1 (lives longest) x2 x3 ← backward use x3 first then x2 x1 last

Follow x1 across the whole picture. It is created first and consumed last, which is why every layer’s activation is still alive at the moment the backward pass begins.

Now draw the memory. It climbs in steps through the forward pass, one step per layer, and nothing comes back down. It reaches its highest point the instant the forward pass finishes, because that is the one moment when every layer’s saved values are alive at once. Then it drains as the backward pass consumes them.

mem FORWARD (L1→L4) BACKWARD (L4→L1) PEAK — every layer held at once +L1+L2+L3+L4 −L4−L3−L2−L1 forward layer ADDS its activations · backward layer CONSUMES then FREES them · peak sets your VRAM need

The curve has exactly one peak, and it sits at the boundary between the forward and backward passes. That peak is the number that has to fit on the card.

Two practical consequences fall straight out of that curve.

First, you cannot free memory between the forward pass and the backward pass. They are not two phases with a gap between them. They share one live set of values, and the handover is the peak.

Second, this gives you a free diagnostic. If a run loads the model fine and then dies a second or two into the first step, that is an activation problem almost every time. Weights run out of memory at load. Activations run out at the forward and backward boundary, so the timing of the crash tells you which tenant to go after.

Turn a tensor’s shape into a number of bytes

Time to put real numbers on all of this. Sizing an activation needs four quantities, and three of them are probably new words.

Token. Text is chopped into tokens before the model sees it. A token is roughly a word or a piece of one. “Fine-tuning” might be two or three tokens.

Sequence length. How many tokens are in one training example. Our running example uses 1,024.

Batch size. How many examples the model processes at the same time. Our running example uses 4. You type this one into a config, which is the first hint that activation memory is under your control in a way weight memory is not.

Hidden size. Inside the model, each token is carried as a list of numbers rather than as a word. The length of that list is the hidden size. For Qwen2.5-1.5B it is 1,536. So a token is not “cat” inside the model, it is 1,536 numbers that encode everything the model currently thinks about that position.

Put those together and a tensor flowing between two operations has shape 4 x 1024 x 1536: four examples, each 1,024 tokens long, each token carrying 1,536 numbers.

Count the elements. 4 * 1024 = 4,096. Then 4,096 * 1,536 = 6,291,456 numbers.

Now bytes per number. This run stores its numbers in bf16, a compact 16-bit floating-point format that uses 2 bytes each instead of the 4 bytes a standard float uses. The next part covers why bf16 is the format everyone reaches for. For now, 2 bytes is all you need.

So 6,291,456 * 2 = 12,582,912 bytes, or about 12.6 MB for one single intermediate tensor. Hold on to that number. We are about to multiply it a great many times.

batchhow many seqs × sequencetokens × hiddenmodel width × bytes2 for bf16 = size

Four numbers set the size of an activation and none of them is the parameter count. Change the batch size or the sequence length and this box changes with them.

Notice what never appeared in that calculation: 1.54 billion. An activation’s size is set by how much data you push through the model, never by how large its weight file is.

Count what one transformer layer leaves behind

Twelve point six megabytes is one tensor. To reach a whole model we need to know how many tensors one layer produces, which means a quick tour of what a layer does.

Qwen2.5-1.5B is 28 identical layers stacked on top of each other, with a lookup table at the bottom that turns each token into its first list of 1,536 numbers, and a small output stage at the top. Each of the 28 layers is a transformer block, and a transformer block is a two-stage workshop. The layer count, the hidden size and the head counts below are all from the Qwen team’s technical report.

Stage one: attention

Every token needs to look at the other tokens, because the meaning of a word depends on its neighbours. Attention is the machinery for that, and it works by giving each token three roles.

The layer converts each token’s 1,536 numbers into three shorter lists. A query says what this token is looking for. A key says what it has to offer. A value is what it hands over if it gets picked. Three tensors, three activations. The query is about the size of the tensor that flowed in. The key and the value are smaller in this particular model, for a reason that comes back at the end of this part.

Then every token’s query is compared against every token’s key, producing one score for each pair of tokens. With 1,024 tokens that is 1,024 x 1,024 scores. The comparison is run 12 times in parallel over different slices of the numbers, so the model can attend to several kinds of relationship at once. Each parallel run is called a head, and this model has 12 query heads.

So the score tensor has shape 4 x 12 x 1024 x 1024. That is batch, then heads, then a full square of token-against-token. Take a good look at that shape, because it is the villain of the second half of this article.

Stage two: the feed-forward network

After attention has mixed information between tokens, each token is processed on its own. The feed-forward network widens each token’s 1,536 numbers out to 8,960, passes them through a simple fixed rule applied to each number in turn, then squeezes them back down to 1,536. That widening is where a transformer does most of its per-token work, and the widened tensor is by far the largest ordinary activation in the layer.

The small parts that still cost bytes

Two more things happen around those stages, and they each leave an intermediate behind. Layer normalisation rescales a tensor so its numbers stay in a workable range instead of drifting to extremes. A residual connection adds a stage’s input back onto its output, so the original signal is never lost no matter how deep the stack gets. Each appears twice per layer, and each emits its own tensor.

Add it all up and a single transformer block emits roughly a dozen tensors of the 4 x 1024 x 1536 size, one fat 4 x 1024 x 8960 tensor, and, depending on how attention is implemented, one 4 x 12 x 1024 x 1024 score tensor.

x input ATTENTION Q, K, V projections scores = QKᵀ (seq×seq) context = scores·V all activations FFN expand hidden → 4–6× activation fn project back the biggest one + norms, residuals →

Count the arrows leaving each stage. The widened feed-forward tensor and the attention score grid are the two that dominate the bill, and everything else is a rounding error next to them.

For the operation-by-operation detail of what each stage computes, the full walkthrough of one transformer block covers it. Here we only care what it leaves behind.

Work out the running example’s 8 GB by hand

Everything is now in place. Three multiplications and we have the answer.

The feed-forward intermediate. Shape 4 x 1024 x 8960. Elements: 4,096 * 8,960 = 36,700,160. Bytes at 2 each: 73,400,320, or about 73 MB.

The hidden-width tensors. We said roughly a dozen at 12.6 MB each. Call it ten to stay conservative: about 126 MB.

The attention scores. Shape 4 x 12 x 1024 x 1024. Elements: 4 * 12 = 48, and 1,024 * 1,024 = 1,048,576, so 48 * 1,048,576 = 50,331,648. Bytes at 2 each: 100,663,296, or about 100 MB.

Activation piece, per layer Shape Elements Bytes in bf16
Feed-forward intermediate 4 x 1024 x 8960 36.7M about 73 MB
Ten hidden-width tensors 10 of 4 x 1024 x 1536 62.9M about 126 MB
Attention scores, if built in full 4 x 12 x 1024 x 1024 50.3M about 100 MB
One layer, total about 300 MB
All 28 layers, the peak about 8 GB

Roughly 300 MB per layer, times 28 layers, gives about 8,400 MB. Call it 8 GB of activations, order of magnitude, for a batch of 4 sequences of 1,024 tokens.

Put that next to the weights, carefully

Qwen2.5-1.5B’s weights in bf16 are 1.54 billion * 2 bytes, so about 3 GB. The activations at this batch and sequence are close to three times the size of the weights themselves. For training that is completely normal.

Be careful how you quote that 3 GB, because this is where memory estimates go wrong in public. Three gigabytes is the weights alone. A full fine-tune also holds one gradient per weight, plus two optimizer statistics and a high-precision master copy of every weight. That is the standard mixed-precision AdamW accounting of 16 bytes per parameter, AdamW being the optimizer almost every language model is trained with. It comes to about 24.6 GB of static state for this model, meaning memory that stays occupied for the whole run. Part 2 derives all sixteen of those bytes.

Activations are not inside that 24.6 GB. Add the roughly 8 GB from the table and a full fine-tune of this model needs about 33 GB. A 32 GB card gives you about 29.8 GiB usable, and every real run wants 2 to 3 GB of headroom on top for memory fragmentation and temporary buffers. So it does not fit, and no amount of optimism closes that gap.

Whenever you meet one of these figures, ask which kind it is. A static-state number and a total are different claims, and the distance between them is the subject of this article.

Name the four dials that set the bill, and spot the quadratic one

Weights depend on one thing: how many parameters the model has. Activations depend on four things, and none of them is the parameter count.

  1. Batch size. Yours to set.
  2. Sequence length. Yours to set.
  3. Hidden size. Fixed by the model you chose.
  4. Layer count. Fixed by the model you chose.

Two of the four are training flags. That is the good news in this whole article: when activations are what is killing you, you can usually fix it tonight without changing models.

WEIGHTS params × bytes set the moment the model loads batch & sequence irrelevant you can’t shrink it without quantizing ACTIVATIONS batch × seq × hidden × layers + attention: batch × heads × seq² batch & seq are yours to tune the memory you can actually control

The left column only changes when you pick a different model. The right column changes when you edit a config file, which is why it is the one you can actually do something about.
Weights, gradients and optimizer state Activations
Size is set by parameter count times bytes per parameter batch times sequence times hidden, times layers
Changes when you choose a different model edit a batch size or a sequence length
Alive for the whole run part of a single step
Shrunk by LoRA, quantisation, a smaller model shorter sequences, smaller batches, checkpointing, flash attention
Typical failure looks like runs out of memory while loading runs out of memory a second into the first step

Some of those words are new. LoRA means freezing most of the weights and training a small patch instead. Quantisation means storing each number in fewer bits. Checkpointing and flash attention each get a section of their own below.

The one activation that does not play fair

Look again at the attention score tensor: 4 x 12 x 1024 x 1024. Sequence length appears in it twice, because it is every token compared against every token. Every other activation has sequence in it once.

That difference is enormous. Double the sequence length from 1,024 to 2,048 and every ordinary activation doubles, which is annoying but survivable. The score tensor goes up four times, from about 100 MB per layer to about 400 MB. Double again to 4,096 and it is about 1.6 GB per layer. Across 28 layers that is roughly 45 GB of attention scores alone, on a card that has 32 GB in total, before a single weight is loaded.

1k · 0.1 GB 2k · 0.4 GB 4k · 1.6 GB 8k · 6.4 GB* ×28 layers = OOM score mem each 2× seq → 4× memory

Each doubling of the sequence length multiplies the score memory by four. The distance between the 1k bar and the 8k bar is the entire reason flash attention was invented.

This is the long-context wall, and it is why “just train on longer examples” is never a free change.

The order to pull the levers in

When a run dies on activations, work down this list. It is ordered by how much relief you get per unit of pain.

  1. Cap the sequence length to what your data actually needs. If 95 percent of your examples are under 600 tokens, training at 2,048 just buys padding, the filler tokens added to make every example the same length. The quadratic term rewards this lever more than any other.
  2. Cut the batch size, then recover the effective batch through gradient accumulation. That means running several small batches one after another, adding their gradients together, and only updating the weights once at the end. Gradients are one number per weight regardless of how many examples produced them, so accumulating costs no extra memory. Activations only ever hold one small batch at a time. Hold the effective batch fixed while you do this, because Goyal et al. showed the learning rate has to scale with that number and not with the small batch in front of it. Part 9 sets it up in a real run.
  3. Turn on gradient checkpointing. The structural lever, and the next section.
  4. Only then reach for a smaller model or a quantised one. Changing the model changes your results. The first three do not.

Cut the peak with gradient checkpointing, and measure the cut

Imagine a long calculation on paper, where you keep every line of working because you will need it later. You run out of paper. The fix is obvious: keep every tenth line, throw the rest away, and when you need line 47 again, start from line 40 and redo seven lines. A little more time, a fraction of the paper.

That is gradient checkpointing, also called activation checkpointing, in full. During the forward pass you keep activations only at a sparse set of positions, called checkpoints, and let the rest go. When the backward pass needs a discarded value, you re-run the forward pass for that short stretch, use the value, and drop it again.

The trade is compute for memory, and the exchange rate is good. You pay roughly one extra forward pass per step, which in practice costs 20 to 30 percent more wall-clock time. In return, Chen and colleagues showed that with checkpoints spaced well, the peak stops being proportional to the number of layers and becomes proportional to the square root of it. For 28 layers, the square root is a little over 5.

mem without · peak = all layers with · low peak + recompute bumps memory O(layers) → O(√layers) · cost ≈ one extra forward (~30% slower)

The sawtooth on the right is recomputation happening on demand. Every tooth is one segment of the forward pass being run a second time to rebuild what was thrown away.

Where the checkpoints go

In practice frameworks put the checkpoints at transformer-block boundaries. It needs no tuning and it captures most of the saving. During the forward pass, only the tensor entering each block is kept. During the backward pass, when the work reaches a block, that block’s forward pass is re-run from its saved input to rebuild its dozen interior tensors, which are used and then released before the next block is touched.

So instead of 28 blocks’ worth of interiors alive at once, you hold 28 small boundary tensors plus one block’s interiors. That is where the order-of-magnitude drop comes from.

kept: C C C only checkpoints (C) held through the forward backward: re-run forward from the last C to rebuild these, use, discard

Only the boundary tensors survive the forward pass. Everything between two checkpoints is rebuilt, used and discarded one segment at a time, so only one segment’s interior is ever live.

One thing checkpointing does not touch: weights, gradients and optimizer state. It is an activations-only lever, so if your problem is the 24.6 GB of static state it will not save you.

And one habit worth forming. Checkpointing costs throughput every step, so leaving it on for a run that already fits is paying 25 percent for nothing. Turn it on when activations are the binding constraint, meaning long sequences, a large batch, or a model near the edge of the card. Turn it off when you have headroom and want speed. It is a lever, not a virtue.

Measuring the saving instead of trusting it

You do not have to take any of this on faith. PyTorch keeps its own high-water mark of GPU memory, and you can read it. This runs one training step with checkpointing off, then one with it on, and prints both peaks.

# torch 2.x
import torch

def peak_gb(model, batch, checkpointing):
    torch.cuda.reset_peak_memory_stats()
    if checkpointing:
        model.gradient_checkpointing_enable()
    model(**batch).loss.backward()
    return torch.cuda.max_memory_allocated() / 1e9

print(f"without checkpointing: {peak_gb(model, batch, False):.1f} GB")
print(f"with checkpointing:    {peak_gb(model, batch, True):.1f} GB")

Line by line, assuming you have never written PyTorch before.

  • import torch pulls in PyTorch, the library that owns the tensors and the GPU allocations.
  • torch.cuda.reset_peak_memory_stats() zeroes a counter. PyTorch tracks the largest amount of GPU memory it has had allocated at once, and this resets that high-water mark so the next measurement starts clean.
  • model.gradient_checkpointing_enable() switches on the recompute behaviour described above. Hugging Face model objects provide this method, so it works on any transformer loaded through that library without your writing the segmenting logic.
  • model(**batch) runs the forward pass. batch is a dictionary of named inputs, and the ** is ordinary Python that spreads a dictionary into keyword arguments. It has nothing to do with PyTorch.
  • .loss pulls the single how-wrong-were-we number out of what the model returned.
  • .backward() runs the backward pass. This is the line where all those stored activations are finally consumed, and it is also the line where an activation-driven OOM almost always fires.
  • torch.cuda.max_memory_allocated() reads the high-water mark back, in bytes. Dividing by 1e9 converts it to gigabytes.
  • The two print lines use Python f-strings. {...:.1f} formats the result to one decimal place.

Expect the second number to be clearly lower than the first, and the run itself to take noticeably longer. Notice also that the two calls run in the order False first and True second. That ordering is not cosmetic, and the exercises at the end of this part ask you to work out why.

Pin your PyTorch and Transformers versions before you build anything on top of this. Memory-reporting helpers and checkpointing entry points have both moved between releases.

Remove the quadratic term entirely with flash attention

Checkpointing shrinks activations by agreeing to recompute them. A second, completely different move is available against the worst offender, and it does better than shrinking the score matrix. It never builds it.

Some hardware background first. A GPU has two very different kinds of memory. There is the large pool, tens of gigabytes of it, which is what people mean by VRAM, and it is comparatively slow to reach. Then there is a tiny scratchpad on the chip itself, measured in tens of kilobytes, which is extremely fast. Standard attention builds the whole 1024 x 1024 score grid in the large pool, which is why it costs about 100 MB per layer here.

Flash attention, which Dao et al. introduced, chops the work into tiles small enough to fit in the on-chip scratchpad. It computes one tile’s contribution to the final answer right there, accumulates it, and discards the tile. Then the next tile. The output is mathematically identical to standard attention. The full score grid simply never exists anywhere.

standard: full seq×seq held in VRAM O(seq²) activation flash: one tile live at a time → O(seq) same math, tiny footprint

Both sides produce the same attention output. The right side never allocates the big square in the middle, and that big square is the entire quadratic term.

The memory effect is the structural one: the sequence-squared term disappears and attention memory becomes linear in sequence length. There is a speed benefit too, because moving less data to and from slow memory is most of what made attention slow.

The two levers do different jobs and are meant to be used together. Flash attention removes the quadratic score activation. Checkpointing shrinks the ordinary per-layer activations across the whole stack. Serious recipes run both: a fused attention backend by default, and checkpointing switched on when the remaining linear activations still will not fit.

One caution about defaults. Modern stacks usually select a fused attention backend, either flash attention itself or PyTorch’s scaled dot product attention, rather than a hand-written score matrix. Confirm which one your run picked rather than assuming, because the attention implementation is a setting on the model and a fallback path can quietly cost you 100 MB per layer.

Explain why inference is cheap, and what the KV cache really is

Everything so far has been about training. Serving a trained model is a different world, and the difference explains why a model that will not fine-tune on a card serves happily on it.

Inference has no backward pass. Nothing computed in the forward pass will ever be needed again, so each layer’s output is freed the moment the next layer has read it. Only one layer’s worth of activations is live at any moment, and that 8 GB peak collapses to a sliver.

There is exactly one deliberate exception. A model generates text one token at a time, and each new token attends to every token before it. That means using each earlier token’s key and value, the two lists from the attention stage, which were computed when that token was first processed. Recomputing them at every step would mean redoing the entire prompt for every single token.

So they are kept. That retained slice is the KV cache. It is not a new species of object. It is an activation that inference chooses to remember because remembering it pays for itself.

generating token by token → layer i freed→ layer i+1 freed→ layer i+2 only one layer alive KV CACHE — K & V of every past token, kept for reuse grows with context length & concurrency — this is what fills serving VRAM, not activations

Every other activation is freed as soon as the next layer has read it. K and V are the deliberate exception, kept because the next token generated will ask for them again.

Its size is straightforward: 2 (one for keys, one for values) times the layer count, times the number of key and value heads, times the width each head works on, times the context length, times bytes per number. Run it for our model. Each head works on a slice of 128 numbers, called the head dimension, and the model has 28 layers and 2 key and value heads. So 2 * 28 * 2 * 128 * 2 bytes = 28,672 bytes, about 28 KB per token. Fill a 256K-token context and you hold roughly 7 GB of cache.

Two engineering responses fall straight out of that formula. Grouped-query attention lets several query heads share one key and value head, which is why this model has 12 query heads but only 2 key and value heads. Each of those serves six query heads, so the cache is six times smaller than it would otherwise be. Separately, paged key and value storage keeps the cache in non-contiguous blocks, so a serving system need not reserve the maximum possible length for every request up front. The part on continuous batching and paged attention covers both.

Place activations among the other tenants, and see why LoRA does not help

Step back and put all five memory tenants in one view.

Tenant In training In inference Scales with
Weights about 3 GB in bf16 about 3 GB in bf16 parameter count
Gradients one per weight none parameter count
Optimizer state the largest static piece none parameter count
Activations about 8 GB, the peak one layer at a time batch, sequence, hidden, layers
KV cache not used about 28 KB per token context length
TRAINING INFERENCE weights gradients — none — optimizer — none — activations freed each layer KV cache — n/a — training holds 4 tenants (16 B/param + activations) · inference holds 2 (weights + KV) → the whole cost gap

Read across the activations row. It is the fattest bar on the training side and a sliver on the inference side, and the whole difference is the backward pass.

Now the row people misread, which decides real hardware purchases. Freezing a weight means the optimizer will never update it. That removes the weight’s gradient and its optimizer state, which is where LoRA’s saving comes from: static state for this model drops from about 24.6 GB to about 4 GB at a typical adapter size. Genuine, large, and worth having.

It does nothing at all to activations, for two reasons. First, the forward pass still runs through every frozen layer, because the frozen layers are what computes the answer. Freezing does not skip a layer, it only stops that layer from changing.

Second, the backward pass still runs back through every frozen layer too. LoRA places a small trainable adapter inside every layer, including the earliest ones, so the gradient message has to travel the whole way down the stack to reach the adapter in layer 1. Travelling through layer 20 means performing layer 20’s backward step, and that step needs layer 20’s stored input. Frozen means “do not update”, and it never means “do not participate”.

So a LoRA run pays roughly the same 8 GB activation bill as a full fine-tune. Its total is about 4 GB of static state plus about 8 GB of activations, so around 12 GB, which fits a 32 GB card with room to spare. The saving is real and it is entirely on the static side. The part on LoRA returns to this where it bites.

Key takeaways

  • An activation is any intermediate value the forward pass computes. Training keeps them because the backward pass needs them, and frees each one the moment it has been used.
  • The reason is one multiplication. A weight’s gradient is the message arriving from above multiplied by that layer’s stored input, so losing the input makes the gradient uncomputable.
  • Activation size is batch times sequence times hidden, summed over the layers, plus a quadratic attention term. Parameter count never appears in that formula.
  • Memory climbs through the forward pass, peaks at the forward and backward boundary, then drains. A crash a second into the first step is an activation problem.
  • For Qwen2.5-1.5B at batch 4 and sequence 1,024, activations come to about 8 GB against 3 GB of bf16 weights. Those 8 GB sit on top of the 24.6 GB of static state, never inside it, and every activation figure is order of magnitude.
  • Gradient checkpointing buys a peak proportional to the square root of the layer count for about 20 to 30 percent more time. Flash attention removes the quadratic term outright. Use both.
  • LoRA, and QLoRA which adds a 4-bit base model on top of it, both pay the full activation bill. Only inference escapes it, because only inference has no backward pass.

You can now

  • Estimate a run’s activation memory from four numbers, per “Work out the running example’s 8 GB by hand”.
  • Explain why a stored forward value is required for a weight’s gradient, per “See why the backward pass cannot throw an activation away”.
  • Diagnose an out-of-memory crash from its timing alone, per “Predict the exact moment in a step when memory peaks”.
  • Pull the right lever first when a run does not fit, per “Name the four dials that set the bill, and spot the quadratic one”.
  • Measure the real saving from gradient checkpointing on your own model, per “Cut the peak with gradient checkpointing, and measure the cut”.
  • Say why LoRA shrinks static state and leaves activations untouched, per “Place activations among the other tenants, and see why LoRA does not help”.

Glossary

Activation
Any intermediate value the forward pass produces on its way from input to output. Activations are held in memory because the backward pass needs them to work out the gradients.
Activation memory
The memory holding the forward pass’s intermediate values until the backward pass consumes them. Its size comes from batch size, sequence length, hidden size and layer count, never from parameter count.
Adapter
A small set of extra trainable weights added beside a frozen model, so you train the adapter and leave the model alone. A LoRA adapter is tens of megabytes against a multi-gigabyte model copy. LoRA explained
Attention
The step where each position in the sequence looks at other positions and mixes in whatever it finds useful. It is what lets a model use context instead of reading each token in isolation. Multi-head attention explained
Attention head
One of several parallel copies of the attention computation, each free to look for a different kind of relationship. Their outputs are joined back together at the end. Multi-head attention explained
Attention scores
The grid of numbers saying how much each position attends to each other position. Its shape is batch by heads by sequence by sequence, so it grows with the square of sequence length, which is the wall flash attention was built to remove. Dao et al., FlashAttention
Autograd
The part of a training framework that records what the forward pass did and then computes every gradient for you. It also decides which intermediate values must be kept, which is why exact activation memory depends on the framework. How a neural network learns
Backpropagation
The procedure that computes a gradient for every weight in one sweep backwards through the model, from the loss at the end to the first layer. It works by applying the chain rule one step at a time. The calculus behind backpropagation
Backward pass
Running backwards from the loss through the model to produce a gradient for every weight. It is backpropagation in practice, and it needs the activations the forward pass stored. How a neural network learns
Batch
A group of examples processed together in one step, so the GPU stays busy and the gradient is averaged over several examples instead of one. Bigger batches give a steadier signal and cost more activation memory. Supervised fine-tuning end to end
bf16
A 16-bit number format with 8 exponent bits and 7 mantissa bits, so it reaches as far as fp32 with much coarser steps. It is the training default because it needs no loss scaling. Number formats for training
Chain rule
The rule for finding the slope through a chain of steps: multiply the slope of each step together. It is what lets one loss at the end of the model tell every weight in every layer how to change. The calculus behind backpropagation
Checkpoint
A saved copy of the model at some point in the run, so you can resume from it, compare it, or ship it. Distinct from gradient checkpointing, which is a memory trick with an unfortunately similar name. Supervised fine-tuning end to end
Context length
The longest sequence a model was built to handle, for example 32K tokens. It caps how long your training examples may be; it is not a target to train at. Qwen2.5-1.5B model card
Effective batch
The number of examples that actually go into one weight update: the per-device batch times the accumulation steps times the number of GPUs. This is the number that matters for training behaviour, rather than the batch that happens to fit. Batch size, accumulation and the effective batch
Feed-forward network
The part of a transformer block that processes each position on its own, widening it to a larger size and squeezing it back. It holds most of a transformer’s weights and produces its largest activation. The transformer feed-forward network
Flash attention
An attention implementation that computes the answer in small tiles inside fast on-chip memory, never building the full score grid. Same result, far less memory, and the term that grew with the square of sequence length becomes linear. Dao et al., FlashAttention
Forward pass
Running data through the model from input to output to get a prediction and a loss. Along the way it produces the activations that the backward pass will need. The complete inference path
Frozen weights
Weights marked as not trainable, so they never receive an update. A frozen weight needs no gradient and no optimizer state, which removes 14 of its 16 bytes, though its activations are still stored. LoRA explained
Full fine-tuning
Updating every weight in the model. It costs 16 bytes per parameter of static state under standard mixed-precision AdamW, produces a complete model copy per task, and forgets the most. Training memory and the 16 bytes per parameter
Gradient
One number per weight saying which way to nudge that weight to make the loss smaller, and how steeply the loss responds. Picture the slope under a ball rolling into a valley. How a neural network learns
Gradient accumulation
Running several small batches, adding their gradients together, and only updating the weights once at the end. You get the steadier signal of a big batch while holding just one small batch in memory. Batch size, accumulation and the effective batch
Gradient checkpointing
Throwing away most stored activations and recomputing them during the backward pass. Peak activation memory drops a long way in exchange for roughly 20 to 30 percent more time. Chen et al., Training Deep Nets with Sublinear Memory Cost
Grouped-query attention
An attention design where several query heads share one set of keys and values, which shrinks the KV cache with little quality loss. The example model has 12 query heads over 2 key-value heads, cutting the cache about sixfold. Ainslie et al., GQA
Hidden size
The width of the list of numbers that flows between layers, written d_model. The example model’s is 1,536, and it multiplies straight into activation memory. Inside one transformer block
Inference
Using a trained model to produce output. There is no backward pass and no optimizer, which is why serving a model costs a fraction of what training it does. The complete inference path
KV cache
The keys and values of past tokens, kept during generation so they are not recomputed for every new token. It exists only at inference; training has no generation loop and therefore no KV cache. Continuous batching and paged attention
Layer
One processing stage inside the model, taking a list of numbers in and handing a transformed list out. The example model is 28 transformer layers deep. Inside one transformer block
Layer norm
A step that rescales the numbers flowing through a layer so they stay in a sensible range, which keeps training stable. Modern LLMs use a cheaper version of it called RMSNorm. Inside one transformer block
LoRA
Low-rank adaptation. Freeze the model and learn a small pair of skinny matrices beside each targeted weight matrix, so about 1 percent of parameters train and the static state for a 1.5B run drops from about 24.6 GB to about 4 GB. Hu et al., LoRA
Loss
One number saying how wrong the model was on this batch. Training is the whole business of making it smaller, and a falling loss on its own proves very little. How a neural network learns
Memory tenant
One of the four things sharing the card during training: weights, gradients, optimizer state and activations. All four are live at the same moment, so peak memory is their sum rather than the largest of them. The four tenants and the 16 bytes per parameter
Micro-batch
The batch that actually fits in memory in one go. Several micro-batches are combined by gradient accumulation so the update behaves like one much larger batch. Batch size, accumulation and the effective batch
Multi-head attention
Attention run as several heads in parallel over different slices of each position’s numbers, then recombined. Grouped-query attention is the memory-saving variant current models use. Multi-head attention explained
OOM
Out of memory, the error you get when a run needs more VRAM than the card has. In training it almost always strikes where the forward pass ends and the backward pass begins, which points straight at activations.
Optimizer
The part of training that turns gradients into actual weight changes. The gradient says which way to move, and the optimizer decides how far. Gradients and optimizers explained
Optimizer state
The numbers an optimizer keeps between steps, such as running averages of past gradients. Under standard mixed-precision AdamW it is 12 of the 16 bytes per parameter, which makes it the largest memory tenant. Rajbhandari et al., ZeRO
Parameter
One of the numbers inside the model that training can change. Parameter and weight mean the same thing here, and a 1.5B model has about 1.5 billion of them. Training memory and the 16 bytes per parameter
QLoRA
LoRA with the frozen model stored in 4 bits instead of 16. It cuts the last remaining cost of simply holding the base weights by about four times, at the price of roughly 40 percent less throughput. Dettmers et al., QLoRA
Quantisation
Storing numbers with fewer bits by mapping them onto a small set of allowed values. It saves memory and gives up some accuracy in return. bitsandbytes documentation
Queries, keys and values
The three sets of numbers attention works from. Each position makes a query saying what it is looking for, a key advertising what it offers, and a value carrying what it passes on when its key is matched. Multi-head attention explained
Residual connection
Adding a layer’s input to its output, so information can skip straight past the layer. It is what makes very deep models trainable at all. Inside one transformer block
SDPA
PyTorch’s built-in scaled dot-product attention, which picks an efficient backend for you including a flash-attention style one. Using it saves installing a separate attention package. The toolchain, bottom to top
Sequence length
How many tokens are in one training example after tokenisation. Activation memory grows in step with it, and the attention part grows with its square.
Tensor
A grid of numbers with any number of dimensions. One number is a scalar, a row of them is a vector, a table is a matrix, and anything past that is still a tensor with more dimensions. How a neural network learns
Token
The unit a language model actually reads and writes: a short piece of text, often a word or part of a word. Every length and cost in training is counted in tokens. The complete inference path
Training step
One cycle of the loop: forward pass, loss, backward pass, optimizer step, then clear the gradients. Everything else in a training script is arrangements around those five moves. Supervised fine-tuning end to end
Transformer
The architecture behind every model in this series: a stack of blocks that alternate attention with a feed-forward network. Vaswani et al., Attention Is All You Need
Transformer block
One repeated unit of the model: attention, then a feed-forward network, with normalisation and residual connections around them. A 28-layer model is 28 of these stacked up. Inside one transformer block
Weight
A single learned number inside the model, used to multiply an input on its way through a layer. Weights are what fine-tuning changes, and the only thing it changes. What actually changes inside the model


Practical exercises

Double the batch and recompute the peak

The worked table in this part gives Qwen2.5-1.5B’s per-layer activation pieces at batch 4, sequence 1024: about 73 MB for the feed-forward intermediate, about 126 MB for roughly ten hidden-width tensors, and about 100 MB for the attention scores, summing to about 300 MB per layer and about 8 GB across 28 layers. Recompute all three pieces and the new total if batch goes from 4 to 8, sequence held at 1024.

See the worked solution (opens in a new tab)

Double the sequence length instead, and watch the trap spring

Now hold batch at 4 and instead double sequence length from 1024 to 2048. Recompute the same three per-layer pieces and the new 28-layer total. Compare how much the total grew here against how much it grew in the batch-doubling exercise, and say which term is responsible for the difference.

See the worked solution (opens in a new tab)

Find the bug in a checkpointing measurement

A colleague uses this part’s peak_gb function to compare checkpointing on and off. In a notebook, they first run peak_gb(model, batch, True) in one cell, see a nicely reduced number, then later run peak_gb(model, batch, False) on the same live model object in a new cell to get the “before” number for their writeup. Both readings come back nearly identical and low. Nothing about the model’s accuracy or training behaviour changed. Diagnose what actually happened, and give the fix.

See the worked solution (opens in a new tab)

Size activations for a 7B run at batch 1

Qwen2.5-7B has hidden size 3,584, feed-forward intermediate width 18,944, 28 layers and 28 query attention heads. Using the same method this part used for the 1.5B example, estimate activation memory in bf16 at batch 1, sequence 2048, with attention scores fully materialised (no flash attention). Then combine that figure with the series’ own static-state numbers for Qwen2.5-7B, about 18 GB for LoRA and about 8 GB for QLoRA, and say which method actually fits a 32 GB card at this batch and sequence, keeping the usual safety margin in mind.

See the worked solution (opens in a new tab)

Frequently asked questions

What is activation memory in LLM training?

It is the memory holding every intermediate value the forward pass computed that the backward pass has not yet used. Its size comes from batch size, sequence length, hidden size and layer count, never from parameter count. It reaches its maximum at the instant the forward pass ends and the backward pass begins.

Why does my training run OOM right after the model loads fine?

Loading only pays for the weights, while the activation pile builds up through the forward pass and peaks when every layer’s stored values are alive at once. An out-of-memory error at that moment is an activation problem, not a weight problem. Cut the sequence length or the batch size, or switch on gradient checkpointing.

How much does gradient checkpointing slow training down?

Commonly 20 to 30 percent, because part of the forward pass is run a second time during the backward pass to rebuild values that were discarded. In exchange, peak activation memory drops from roughly proportional to the layer count to roughly proportional to its square root. That trade is often what turns an impossible run into a working one.

Does flash attention save memory or only time?

Both, and the memory effect is the structural one. By computing attention in small tiles inside the GPU’s fast on-chip scratchpad, it never builds the full token-by-token score grid, so the sequence-squared term disappears from the memory bill. That is what makes training on long sequences practical at all.

Why does LoRA not reduce activation memory?

Because a frozen layer still runs in the forward pass and still passes the gradient message back through in the backward pass, so that adapters lower in the stack can be reached. Freezing removes a weight’s gradient and its optimizer state, which is where LoRA’s static-state saving comes from. It does not remove the need to store that layer’s forward input.

Is the KV cache just an activation?

Yes, it is an activation that inference deliberately holds on to. The keys and values of earlier tokens would otherwise be recomputed at every single generation step, so keeping them pays for itself immediately. Its size is 2 times layers times key and value heads times head dimension times context length times bytes per number.

Sources and further reading

Previous