Fine-Tuning LLMs

Course Content

Fine-Tuning LLMs

6 sections · 52 lessons

Why does fine-tuning use more memory than inference (activations, gradients, optimizer states)?


Full fine-tuning an 8B model with mixed-precision AdamWeights,bf16: 16 GBGradients,bf16: 16 GBfp32 masterweights: 32 GBAdam m and v,fp32: 64 GBActivations andlogits: variestopbottomInference needs only the first block plus about 128 KB of KV cache per token.
The optimizer, not the model, is the biggest item — Adam's two fp32 states alone are four times the size of the weights.

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):

Text
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 sequence

Training memory, full fine-tuning

ItemBytes per parameter8B model
Weights (bf16)216 GB
Gradients (bf16)216 GB
fp32 master weights432 GB
Adam momentum and variance (fp32)864 GB
Subtotal16about 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

LeverSavesCost
LoRA / QLoRAGradients and optimizer states for frozen weights; QLoRA also shrinks weights to 4-bitSlight quality risk; QLoRA slower
Gradient checkpointingMost activations — recomputed in the backward passAbout one extra forward pass per step
Smaller micro-batch + gradient accumulationActivationsMore steps per update
8-bit or paged optimizersAdam's 8 bytes drop to about 2Tiny quality risk
FSDP / DeepSpeed ZeROSplits weights, gradients and optimizer states across GPUs (ZeRO-3 on 8 GPUs: about 128 / 8 = 16 GB each)Communication between GPUs
CPU offloadMoves optimizer states to CPU RAMMuch 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.