How to Handle Out-of-Memory Errors When Training Large Models on Limited GPU Memory
You can prevent out-of-memory (OOM) errors when training large transformers on limited GPU memory by reducing batch size and context length, implementing gradient accumulation, enabling mixed-precision training with Automatic Mixed Precision (AMP), and applying activation checkpointing to recompute intermediate activations during the backward pass.
Training transformer models with millions or billions of parameters quickly exhausts GPU memory when using default configurations like batch size 32, context length 512, and full-precision (FP32) tensors. The FareedKhan-dev/train-llm-from-scratch repository provides a straightforward training loop in scripts/train_transformer.py that can be adapted using several memory optimization techniques to stay within constrained hardware budgets.
Reduce Batch Size and Context Length
The simplest way to decrease memory consumption is to reduce the number of tokens processed simultaneously. In the train-llm-from-scratch repository, these values are controlled via config/config.py.
Open config/config.py and modify these variables:
# config/config.py
T_BATCH_SIZE = 8 # Default might be 32; reduce to fit your GPU
T_CONTEXT_LENGTH = 256 # Default might be 512; shorter sequences use less memory
Smaller batches mean fewer activations stored per step, while shorter sequences lower the size of activation tensors in each transformer block. This approach trades training speed for memory efficiency without requiring code changes to the model architecture.
Implement Gradient Accumulation
Gradient accumulation simulates a larger effective batch size by accumulating gradients over multiple forward-backward passes before calling optimizer.step(). This technique allows you to train with the same convergence properties as large batches while keeping the per-step memory footprint small.
Add an accumulation step configuration to config/config.py:
# config/config.py
ACCUM_STEPS = 4 # Accumulate gradients over 4 mini-batches
Then modify the training loop in scripts/train_transformer.py to implement accumulation logic:
# scripts/train_transformer.py (excerpt)
accum_counter = 0
optimizer.zero_grad(set_to_none=True)
for step, (xb, yb) in enumerate(batch_iterator):
# Forward pass
_, loss = model(xb, yb)
# Scale loss by accumulation steps to maintain gradient magnitude
(loss / config['accum_steps']).backward()
accum_counter += 1
# Only update weights after accumulating enough gradients
if accum_counter % config['accum_steps'] == 0:
optimizer.step()
optimizer.zero_grad(set_to_none=True)
accum_counter = 0
Setting set_to_none=True in zero_grad() is more memory-efficient than setting gradients to zero, as it allows Python to garbage collect the memory immediately.
Enable Mixed-Precision Training (AMP)
Automatic Mixed Precision (AMP) stores tensors in FP16 where possible, cutting memory usage roughly in half while maintaining training stability through loss scaling. This is implemented using torch.cuda.amp.autocast and torch.cuda.amp.GradScaler.
Modify scripts/train_transformer.py to wrap the forward-backward pass:
from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
for step, (xb, yb) in enumerate(batch_iterator):
optimizer.zero_grad(set_to_none=True)
with autocast():
_, loss = model(xb, yb)
# Scale loss to prevent gradient underflow in FP16
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
The autocast context manager automatically selects FP16 for matrix multiplications where safe, while GradScaler scales the loss value to prevent vanishing gradients in the backward pass.
Apply Activation Checkpointing
Activation checkpointing (also called gradient checkpointing) trades compute for memory by recomputing intermediate activations during the backward pass instead of storing them. This dramatically reduces memory usage, especially in deep transformer models with many layers.
In the repository, modify src/models/transformer_block.py to apply checkpointing to each transformer block:
import torch.utils.checkpoint as checkpoint
class Block(nn.Module):
def forward(self, x):
def block_forward(x):
# Original block computation: attention + MLP
return self._block_impl(x)
# Use checkpointing to avoid storing intermediate activations
return checkpoint.checkpoint(block_forward, x)
Apply torch.utils.checkpoint.checkpoint to the forward pass of each block to enable this optimization. The memory savings scale with model depth, though each backward pass will require additional compute to recompute the activations.
Clear Cached Memory After Evaluation
During evaluation phases, PyTorch may retain memory caches that aren't immediately needed for training. Force the allocator to release unused memory by calling torch.cuda.empty_cache() after evaluation functions.
Update the estimate_loss function in scripts/train_transformer.py:
def estimate_loss(steps: int) -> Dict[str, float]:
# ... existing evaluation logic ...
# Release unused cached memory back to the GPU
torch.cuda.empty_cache()
return out
This is particularly useful when alternating between training and validation phases where the validation batch size might differ significantly from training.
Summary
Training large language models on limited GPU memory requires strategic trade-offs between batch size, numerical precision, and computational overhead:
- Reduce configuration values (
T_BATCH_SIZEandT_CONTEXT_LENGTHinconfig/config.py) to immediately lower memory requirements without code complexity - Accumulate gradients over multiple steps to maintain large effective batch sizes while keeping per-step memory constant
- Use AMP (
torch.cuda.amp) to halve tensor memory usage through FP16 precision with automatic loss scaling - Checkpoint activations in
src/models/transformer_block.pyto reduce memory linearly with model depth at the cost of additional forward computations - Clear GPU caches after evaluation phases with
torch.cuda.empty_cache()to reclaim unused allocated memory
Frequently Asked Questions
What is the most effective single change to prevent OOM errors when training transformers?
Reducing the batch size in config/config.py is the most immediate and effective change, as it linearly reduces the number of activation tensors stored during the forward pass. If you must maintain a specific effective batch size for convergence, combine a smaller T_BATCH_SIZE with gradient accumulation steps to achieve the same result without the memory penalty.
Does mixed-precision training affect model accuracy or convergence?
Mixed-precision training with torch.cuda.amp typically maintains model accuracy while using approximately half the memory of FP32 training. The GradScaler component automatically handles loss scaling to prevent gradient underflow, and modern GPUs (Tensor Cores) actually accelerate FP16 operations, often resulting in faster training times alongside memory savings.
When should I use activation checkpointing versus simply reducing batch size?
Use activation checkpointing when you have already minimized batch size to 1 and still encounter OOM errors, or when you need to train with longer context lengths that exceed memory limits. Checkpointing is particularly effective for deep models (many transformer blocks) where memory usage scales with depth, whereas batch size reduction helps when sequence length or width causes the OOM.
How do I determine the optimal number of gradient accumulation steps?
Set accumulation steps based on your target effective batch size divided by the maximum batch size that fits in memory. For example, if you need an effective batch of 32 but can only fit batch size 4 in memory, set ACCUM_STEPS = 8 in config/config.py. Monitor GPU memory usage with nvidia-smi or torch.cuda.memory_allocated() to ensure you are utilizing available memory without exceeding limits.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →