Most explanations of LoRA stop at the algebra: freeze the weight matrix, learn a low-rank correction, save a tiny file. That is true and it hides the questions you actually face when a training job is about to start on a real GPU. How much memory will this step need? Why is the job only somewhat faster than full fine-tuning when it trains half a percent of the parameters? Why did memory spike at the loss, and why did enabling gradient checkpointing silently stop the adapters from learning?
This article follows one LoRA training step through the GPU for an 8B decoder model with the Llama-3-8B shape. It counts FLOPs and bytes from first principles, shows where activations and the logits tensor dominate, explains what checkpointing and sharding buy you, and ends with a plan for a 24 GB class card and an 80 GB class card. Every figure is arithmetic you can redo; activation memory is something you should measure, and the code shows how.
What LoRA changes inside a training step
A linear layer computes h = W0 x. LoRA keeps W0 frozen and adds a parallel path: h = W0 x + s * B (A x), where A has shape r by in, B has shape out by r, and the scale s is alpha / r (or alpha / sqrt(r) with rank-stabilised LoRA). B starts at zero, so step zero reproduces the base model exactly.
For the GPU the important consequences are about the backward pass. A training step for a dense matrix has three matrix multiplies: the forward product, the input gradient dx = W0^T dh that earlier layers need, and the weight gradient dW0 = dh x^T. Freezing W0 removes only the third. The input gradient still has to flow through every frozen layer, because the adapters in layer 3 need the gradient that arrives from layer 30. So LoRA does not make the backward pass free; it makes it roughly half as expensive.
The second consequence is about state. The frozen weights need no gradient buffer and no optimizer state. Adam keeps two moments per trainable parameter, so removing them from 8 billion parameters is where the large memory saving comes from. The original LoRA paper reported that for GPT-3 175B, GPU memory during training fell by about a factor of three and the number of trainable parameters by about ten thousand.
Compute per token: 4N, not 6N, and what erodes it
A useful rule for transformer training is about 6N matrix FLOPs per token, where N is the parameter count: 2N forward, 2N for input gradients and 2N for weight gradients. With frozen base weights it is about 4N, because the weight-gradient term disappears. For an 8.03B model that is roughly 48 GFLOPs per token for full fine-tuning against 32 GFLOPs for LoRA. The adapters themselves add 2 * r * (in + out) FLOPs per module per token in the forward pass; at rank 16 across all linear layers that is about 84 MFLOPs, a fraction of one percent.
Two things erode the ideal one-and-a-half times speedup. Attention score and value computations are not counted in 6N and are not reduced by LoRA at all, so at long sequence lengths their share grows. And gradient checkpointing, which nearly every memory-constrained LoRA run uses, recomputes the forward pass during backward, adding another 2N: LoRA plus checkpointing is about 6N, the same as full fine-tuning without it. In practice the win from LoRA is that the job fits on far fewer GPUs, not that each token is much cheaper.
The memory budget for an 8B model
Here is the budget for the Llama-3-8B shape (hidden 4096, key and value width 1024, MLP width 14336, 32 layers, vocabulary 128256) with rank 16 on all seven linear modules per layer. The trainable count is 32 * 16 * (2*(4096+4096) + 2*(4096+1024) + 3*(4096+14336)) = 41,943,040, which is what PEFT prints for this configuration.
| Item | LoRA r=16, bf16 base | Full fine-tune, mixed precision Adam |
|---|---|---|
| Base weights | 16.1 GB (bf16) | 16.1 GB (bf16) plus 32 GB fp32 master copy |
| Weight gradients | none for the base; 0.17 GB fp32 for adapters | 16.1 GB (bf16) |
| Adam moments | 0.34 GB for adapters | 64 GB |
| Adapter weights | 0.17 GB fp32 | not applicable |
| Activations | scales with tokens per micro-batch; measure | same order as LoRA |
| Logits and loss | about 2.1 GB fp32 per 4,096-token sequence, plus its gradient | same |
Adapter state totals 671 MB, which is noise. What remains is the 16 GB of frozen weights and everything that scales with tokens. Activations are the term people underestimate: each layer saves the inputs to its matrix multiplies, the attention softmax statistics, the MLP intermediate (14336 wide, three times the hidden width) and normalisation inputs. LoRA adds to this, because each adapter must save x and A x for its own backward, and if lora_dropout is nonzero the dropped copy of x is stored as well.
The logits are a separate spike. For one 4,096-token sequence the output projection produces 4,096 by 128,256 values; at fp32, which many Hugging Face model versions upcast to before the loss, that is 2.1 GB, and the softmax gradient is another tensor of the same size. With four sequences per micro-batch the loss alone wants well over 8 GB. Chunked or fused cross-entropy kernels that never materialise the full logits are one of the most effective memory fixes for large-vocabulary models.
The step in one picture
Activations, checkpointing and a training loop that measures itself
Gradient checkpointing keeps only each decoder layer's input and recomputes the layer during backward. For our model the saved inputs for one 4,096-token sequence are 32 * 4096 * 4096 * 2 bytes, about 1.07 GB in bf16, plus the working set of whichever single layer is being recomputed. That turns activation memory from the dominant term into a modest one, at the cost of the extra forward pass discussed above.
There is a well-known trap. With reentrant checkpointing (the older use_reentrant=True implementation), a checkpointed segment produces outputs that require gradients only if one of its inputs does. In a LoRA model the embedding layer is frozen, so the hidden states entering layer 0 do not require gradients, the checkpointed layers report no gradient path, and the adapters receive no gradient. PyTorch warns that none of the inputs have requires_grad=True. Then either loss.backward() fails because the loss has no gradient function, or, if something else such as a head is trainable, every adapter gets None gradients and the loss barely moves. The fixes are to call model.enable_input_require_grads() (which PEFT's preparation helpers do), or to use non-reentrant checkpointing.
import torch
from transformers import AutoModelForCausalLM
from peft import LoraConfig, get_peft_model
model = AutoModelForCausalLM.from_pretrained(
MODEL_ID, torch_dtype=torch.bfloat16, attn_implementation="sdpa").cuda()
model.gradient_checkpointing_enable(
gradient_checkpointing_kwargs={"use_reentrant": False})
model.config.use_cache = False # the KV cache is useless in training
cfg = LoraConfig(r=16, lora_alpha=32, lora_dropout=0.0,
target_modules="all-linear", task_type="CAUSAL_LM")
model = get_peft_model(model, cfg)
model.print_trainable_parameters() # expect 41,943,040 for this shape
params = [p for p in model.parameters() if p.requires_grad]
print({p.dtype for p in params}) # adapter dtype depends on your PEFT version
opt = torch.optim.AdamW(params, lr=2e-4, weight_decay=0.0)
for step, batch in enumerate(loader): # packed sequences, labels masked with -100
torch.cuda.reset_peak_memory_stats()
out = model(**{k: v.cuda() for k, v in batch.items()})
out.loss.backward()
if step == 0: # every adapter needs a grad; A's is exactly 0 while B is 0
dead = [n for n, p in model.named_parameters() if p.requires_grad and
(p.grad is None or ("lora_B" in n and not p.grad.any()))]
assert not dead, f"{len(dead)} adapter tensors got no gradient: {dead[:3]}"
torch.nn.utils.clip_grad_norm_(params, 1.0)
opt.step(); opt.zero_grad(set_to_none=True)
if step < 3 or step % 50 == 0:
gb = torch.cuda.max_memory_allocated() / 2**30
print(f"step {step} loss {out.loss.item():.4f} peak {gb:.2f} GiB")Print the peak for the first few steps at the sequence length and micro-batch you intend to use; the first step includes allocator warm-up and the optimizer state appears after the first opt.step(). The assertion after the first backward is the cheap guard against the checkpointing trap: every adapter tensor must have a .grad, and the B matrices must have nonzero ones. Do not demand nonzero gradients for A on step zero: the gradient of A is proportional to B, which is still exactly zero.
Small GEMMs and throughput
With rank 16, the adapter GEMMs are thin: the A x product multiplies a tokens-by-4096 matrix by a 4096-by-16 matrix. Its arithmetic intensity is low, so it is bound by memory bandwidth rather than tensor-core throughput, and each one is a separate kernel launch. Seven modules, two products each, forward and backward, across 32 layers is well over a thousand small kernels per step that a full fine-tune does not run, visible in a profiler as a long tail of tiny kernels.
Practical mitigations, roughly in order of effort: raise tokens per micro-batch by packing sequences instead of padding them; set lora_dropout=0 so no extra copy of the input is made; try torch.compile on the model, which can fuse the scale and add into neighbouring kernels; and consider libraries that ship fused LoRA kernels if the profile shows the adapter path is a large share of step time. Measure before and after with torch.profiler; do not assume a gain that a vendor benchmark reports on a different shape.
Scaling out: DDP, FSDP and a 4-bit base
Data parallelism is cheap with LoRA because only the adapter gradients are all-reduced: 42 million values, about 168 MB in fp32, against 16 GB of bf16 gradients for a full fine-tune. If one GPU holds the frozen model, plain DDP with one replica per GPU is the simplest way to scale.
If the frozen weights do not fit, shard them. With PyTorch FSDP1, one wrapped unit flattens its parameters into a single buffer, and mixing frozen and trainable parameters in one unit requires use_orig_params=True; FSDP2's per-parameter fully_shard does not have that restriction. DeepSpeed ZeRO stage 3 is the equivalent alternative. The other route is to shrink the frozen weights to 4 bits, which is QLoRA: it trades some step time for dequantisation for a base that takes about a quarter of the memory.
Worked example: planning for two GPU sizes
Suppose you want to adapt the 8B model on support conversations with sequences packed to 4,096 tokens.
- 24 GB class card, bf16 base. Weights 16.1 GB, adapter state 0.67 GB, checkpointed layer inputs 1.07 GB per sequence, fp32 logits and their gradient about 4.2 GB per sequence, plus one layer's recompute working set and allocator overhead. One sequence per micro-batch is at the edge and may not fit; this is the case for QLoRA or a fused loss kernel, with gradient accumulation to reach the effective batch.
- 80 GB class card, bf16 base. About 60 GB is left after weights. Start with four sequences per micro-batch with checkpointing, measure the peak, and increase until the peak is about 90 percent of capacity. If the profile shows the loss dominating memory, fix that before turning off checkpointing.
- Several GPUs. If the model fits per GPU, use DDP; the communication is tiny. Shard only when it does not fit, because sharding adds all-gathers of the frozen weights in every forward and backward.
Failure modes
- Backward fails, or loss flat from step one. Reentrant checkpointing without input gradients (an error at backward, or silently None adapter gradients, depending on what else is trainable), adapters attached to the wrong module names, or an optimizer built before
get_peft_modelthat holds no adapter parameters. - Out of memory at the loss, not in the layers. The fp32 logits spike. Reduce micro-batch, chunk the loss, or use a fused cross-entropy.
- Slower than expected. Many tiny adapter kernels on short padded sequences. Pack, raise tokens per step, profile.
- NaNs with fp16. Use bf16 on hardware that supports it; fp16 needs loss scaling. Keep adapter weights in fp32.
- FSDP errors about mixed requires_grad. FSDP1 flat parameters; set
use_orig_params=Trueor wrap per module. - Saved adapter does not reproduce evaluation. Different scale convention (rsLoRA against plain), or merged into a quantized base. Merge into the precision you trained against and re-evaluate.
Trade-offs
Rank trades adapter capacity for compute and state, but at the ranks people use the GPU cost of the adapter is small; the base model and activations decide what fits. Targeting all linear layers usually learns more than attention-only targets for the same rank, with more small kernels. Checkpointing trades extra compute for a large cut in activation memory: about half again for LoRA (4N to 6N per token), against about a third for full fine-tuning (6N to 8N). A 4-bit base trades step time and a little fidelity for fitting on one card. Full fine-tuning remains the right tool when the change is large and you have the memory, because it updates every weight and runs fewer kernels per token.
What to do next
- Compute your adapter count with the formula above and confirm it against
print_trainable_parameters(). - Run three steps at your real sequence length and print
max_memory_allocated; record it per micro-batch size. - Assert that every adapter parameter has a nonzero gradient after the first backward.
- Profile one step with
torch.profilerand note the share of time in the loss and in small adapter kernels; see GPU profiling. - Pack sequences, set dropout to zero unless validation needs it, and retest throughput.
- If the frozen model does not fit, choose between FSDP sharding and a QLoRA 4-bit base.
- Review mixed-precision training for dtype choices, and the small-model LoRA workflow for data, rank and evaluation.