Course Content
Fine-Tuning LLMs
6 sections · 52 lessons
Why does fine-tuning use more memory than inference (activations, gradients, optimizer states)?
What you need to know
Inference memory
Weights plus the KV cache — the keys and values of past tokens that attention re-reads. For Llama-3.1-8B (32 layers, 8 key/value heads of size 128, bf16):
weights: 8.03B x 2 bytes = about 16 GBKV cache: 2 (K and V) x 32 x 8 x 128 x 2 bytes = 128 KB per token x 4,096 tokens = 512 MB per sequenceTraining memory, full fine-tuning
| Item | Bytes per parameter | 8B model |
|---|---|---|
| Weights (bf16) | 2 | 16 GB |
| Gradients (bf16) | 2 | 16 GB |
| fp32 master weights | 4 | 32 GB |
| Adam momentum and variance (fp32) | 8 | 64 GB |
| Subtotal | 16 | about 128 GB |
Then add activations: every layer's intermediate results, kept from the forward pass so the backward pass can compute gradients. They grow with batch size × sequence length × layers, and for long sequences they can exceed everything else.
One often-missed item is the logits: the model's scores over the whole vocabulary for every token. With a 128,256-token vocabulary and 32,768 tokens in a batch, the fp32 logits alone are 32,768 × 128,256 × 4 bytes ≈ 16.8 GB. That is why TRL 1.13's SFTConfig defaults to a chunked loss, which never builds the whole logits tensor at once.
Levers when you run out of memory
| Lever | Saves | Cost |
|---|---|---|
| LoRA / QLoRA | Gradients and optimizer states for frozen weights; QLoRA also shrinks weights to 4-bit | Slight quality risk; QLoRA slower |
| Gradient checkpointing | Most activations — recomputed in the backward pass | About one extra forward pass per step |
| Smaller micro-batch + gradient accumulation | Activations | More steps per update |
| 8-bit or paged optimizers | Adam's 8 bytes drop to about 2 | Tiny quality risk |
| FSDP / DeepSpeed ZeRO | Splits weights, gradients and optimizer states across GPUs (ZeRO-3 on 8 GPUs: about 128 / 8 = 16 GB each) | Communication between GPUs |
| CPU offload | Moves optimizer states to CPU RAM | Much slower steps |
A real-life example
The radiology team trains a rank-64 LoRA on an 8B model on one 48 GB GPU: micro-batch 8, reports up to 4,096 tokens. It crashes with out-of-memory on step 1.
Their estimate: base weights 16 GB and the adapter's weights, gradients and Adam states under 3 GB — so activations and logits are the problem. At 8 × 4,096 = 32,768 tokens per step, the logits alone are about 16.8 GB in fp32. They turn on gradient checkpointing, cut the micro-batch to 2 with 4 steps of gradient accumulation (the same effective batch of 8), and keep TRL's chunked loss. Peak memory falls to about 30 GB, and each update is about 25% slower.
Follow-up questions to expect
- "Why keep fp32 master weights?" — Updates are tiny. In bf16, 1.0 plus 0.001 rounds back to 1.0, so small updates would be lost. The fp32 copy accumulates them.
- "Does LoRA reduce activation memory?" — Only slightly. Activations flow through the full frozen network; use checkpointing and shorter micro-batches.
- "Why does inference also need a lot of memory at scale?" — The KV cache grows with users × context length; at 100 concurrent 4,096-token sequences, the example above needs about 50 GB of KV cache on top of the weights.