How to Optimize Transformer Training Performance on a Single GPU: 8 Proven Techniques
Enable mixed-precision training with torch.cuda.amp, replace the manual attention computation in src/models/attention.py with torch.nn.functional.scaled_dot_product_attention, and use a DataLoader with pin_memory=True to eliminate CPU-GPU transfer bottlenecks, yielding up to 2.5× higher throughput on the train-llm-from-scratch repository.
Training large language models on a single GPU quickly exhausts memory bandwidth and compute resources due to kernel launch overhead and inefficient data movement. The train-llm-from-scratch repository by FareedKhan-dev provides a clean, modular GPT-style transformer implementation, but its default configuration leaves significant performance gains unclaimed. This guide identifies eight architecture-specific optimizations—grounded in the exact source files—that maximize tokens-per-second without altering model correctness or mathematical behavior.
Enable CUDA Backend Optimizations
Two one-line configuration changes in scripts/train_transformer.py unlock immediate speedups by tuning the CUDA backend for transformer workloads.
CuDNN Autotuning configures the cuDNN library to profile convolution and softmax kernels during the first few training iterations, then select the fastest algorithms for your fixed context length. TF32 Precision allows Ampere-generation GPUs (RTX 30-series/40-series and A100) to utilize TensorFloat-32 matrix cores, doubling matmul throughput with negligible accuracy loss compared to FP32.
Add these directives immediately after your imports:
import torch
torch.backends.cudnn.benchmark = True # Profile and select fastest kernels
torch.backends.cuda.matmul.allow_tf32 = True # Enable TF32 for matmul on Ampere+
Implement Mixed-Precision Training with Gradient Accumulation
Modern NVIDIA GPUs execute FP16 or BF16 operations at twice the rate of FP32 while consuming half the memory bandwidth. Wrapping the forward-backward pass in torch.cuda.amp.autocast() and using a GradScaler enables automatic loss scaling to prevent gradient underflow. When the desired batch size exceeds available VRAM, gradient accumulation maintains the effective batch size by splitting the workload into micro-batches and only stepping the optimizer after accumulation completes.
Modify the training loop in scripts/train_transformer.py to implement four-step accumulation:
import torch.cuda.amp as amp
ACCUM_STEPS = 4
scaler = amp.GradScaler()
for step in pbar:
xb, yb = next(batch_iterator)
with amp.autocast():
_, loss = model(xb, yb)
loss = loss / ACCUM_STEPS
scaler.scale(loss).backward()
if (step + 1) % ACCUM_STEPS == 0:
scaler.unscale_(optimizer)
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad(set_to_none=True)
Replace Naive Attention with Flash-Attention Kernels
The Head class in src/models/attention.py currently implements scaled dot-product attention via explicit matrix multiplication followed by manual masking and softmax. This approach is memory-bound and underutilizes the GPU's tensor cores. Replacing this with PyTorch 2.0's fused scaled_dot_product_attention kernel reduces memory traffic by 2-3× and automatically selects the most efficient implementation, including FlashAttention algorithms when available.
Update the forward method in src/models/attention.py as follows:
def forward(self, x: torch.Tensor) -> torch.Tensor:
B, T, C = x.shape
k = self.key(x)
q = self.query(x)
v = self.value(x)
attn = torch.nn.functional.scaled_dot_product_attention(
q, k, v,
attn_mask=self.tril[:T, :T].bool(),
dropout_p=0.0,
is_causal=True
)
return attn
Compile the Model Graph with torch.compile
Python interpreter overhead and operator fragmentation create significant latency in the standard eager execution mode. Adding model = torch.compile(model) immediately after model construction in scripts/train_transformer.py JIT-compiles the entire transformer into optimized CUDA graphs. This fusion eliminates redundant kernel launches and reduces Python GIL contention, typically yielding 10-20% additional throughput on top of other optimizations.
model = Transformer(config)
model = torch.compile(model) # JIT compilation for optimized execution
model = model.to(device)
Eliminate Data Loading Bottlenecks with Pinned Memory
The default implementation in data_loader/data_loader.py uses a synchronous generator that blocks the GPU while copying batches from CPU memory. Replacing this with a torch.utils.data.IterableDataset wrapped in a standard DataLoader configured with pin_memory=True enables asynchronous CPU-to-GPU transfers. This ensures the CUDA compute units remain saturated by prefetching data into page-locked memory.
Refactor data_loader/data_loader.py to use an iterable dataset:
from torch.utils.data import IterableDataset, DataLoader
import h5py
import numpy as np
import torch
class H5IterableDataset(IterableDataset):
def __init__(self, data_path, context_length):
self.path = data_path
self.context = context_length
def __iter__(self):
with h5py.File(self.path, 'r') as h5:
tokens = h5['tokens']
N = tokens.shape[0] // (self.context + 1)
order = np.random.permutation(N)
for idx in order:
start = idx * self.context
seq = tokens[start:start + self.context + 1]
xb = torch.tensor(seq[:-1], dtype=torch.long)
yb = torch.tensor(seq[1:], dtype=torch.long)
yield xb, yb
Then instantiate in scripts/train_transformer.py with pinned memory:
train_dataset = H5IterableDataset(config['train_path'], config['t_context_length'])
batch_iterator = DataLoader(
train_dataset,
batch_size=config['t_batch_size'],
pin_memory=True, # Enables async CPU->GPU transfer
num_workers=2, # Parallel prefetching
drop_last=True
).__iter__()
Reduce Evaluation Overhead
The estimate_loss function runs a full forward pass over the validation set every t_eval_steps iterations as defined in config/config.py. These validation runs consume significant compute cycles that could contribute to training. Increasing t_eval_steps or reducing t_eval_iters minimizes this overhead without compromising training quality or convergence monitoring.
Adjust these parameters in config/config.py:
config = {
't_eval_steps': 500, # Increase from default to evaluate less frequently
't_eval_iters': 20, # Reduce number of validation batches per evaluation
}
Summary
- Enable backend optimizations by setting
cudnn.benchmarkandmatmul.allow_tf32inscripts/train_transformer.pyto unlock hardware-accelerated kernels. - Use mixed-precision training with
torch.cuda.amp.autocastandGradScaler, combined with gradient accumulation to simulate larger batch sizes within existing memory constraints. - Replace manual attention in
src/models/attention.pywithtorch.nn.functional.scaled_dot_product_attentionto leverage fused FlashAttention kernels. - JIT-compile the model using
torch.compileimmediately after instantiation to reduce Python overhead. - Optimize data loading by converting the custom generator in
data_loader/data_loader.pyto aDataLoaderwithpin_memory=Truefor asynchronous transfers. - Tune evaluation frequency in
config/config.pyto reduce intermittent validation bottlenecks.
Frequently Asked Questions
Does mixed-precision training affect model convergence or final loss?
No. Using torch.cuda.amp with automatic loss scaling maintains numerical stability equivalent to full FP32 training. The GradScaler dynamically adjusts the loss scale factor to prevent gradient underflow in the backward pass, ensuring the optimization trajectory and final validation loss remain statistically identical while delivering 1.5-2× speedups on modern NVIDIA GPUs.
Can I apply these optimizations to older GPUs that do not support TF32?
Yes. While torch.backends.cuda.matmul.allow_tf32 requires Ampere architecture or newer, mixed-precision training via torch.cuda.amp.autocast is supported on Pascal (SM 6.0) and newer GPUs. The scaled_dot_product_attention function automatically detects hardware capabilities and falls back to memory-efficient attention implementations if FlashAttention kernels are unavailable, though throughput gains will be most significant on RTX 30-series or newer hardware.
How does gradient accumulation interact with the learning rate schedule?
Gradient accumulation preserves the effective batch size while decoupling it from the micro-batch size processed per forward pass. You should configure your learning rate based on the effective batch size (micro-batch size × ACCUM_STEPS), not the micro-batch size alone. The optimizer step only occurs after accumulation completes, so the learning rate schedule should count optimizer steps rather than forward passes to maintain consistent convergence behavior.
Where exactly should I place the torch.compile call in the training script?
Place model = torch.compile(model) immediately after model construction in scripts/train_transformer.py and before moving the model to the GPU with .to(device). According to the PyTorch 2.0 documentation, compiling before device placement ensures the graph capture includes the correct device metadata, though torch.compile handles device semantics automatically in most cases. Avoid compiling inside the training loop to prevent re-compilation overhead.
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 →