Why Scale Matters
Training a small logistic regression on a laptop takes seconds. Training a modern large language model takes weeks across thousands of GPUs. Between these extremes lies a spectrum of engineering challenges — memory limits, communication overhead, numerical precision, and hardware utilization — that determine whether a training run is feasible, affordable, and fast.
Understanding how to scale training efficiently is no longer optional for ML practitioners. Even modest workloads — a ResNet on ImageNet, a transformer for text classification — benefit enormously from the techniques covered here. The core ideas are general: distribute computation, exploit hardware efficiency, and reduce memory footprint without sacrificing accuracy.
Hardware: GPU and TPU
Modern deep learning is dominated by massively parallel matrix operations — exactly the workload that GPUs were designed to accelerate. A GPU contains thousands of smaller cores (CUDA cores on NVIDIA, shader units on AMD) that execute the same instruction on many data elements simultaneously (SIMD). For matrix multiplications — the backbone of neural network forward and backward passes — GPUs can be 10–100× faster than CPUs.
GPU Memory Hierarchy
GPU VRAM (video RAM) is the on-chip memory where model parameters, activations, and gradients must reside during training. Modern GPUs carry 16–80 GB of VRAM (e.g., NVIDIA H100: 80 GB HBM3). The GPU communicates with the CPU host over PCIe, which is orders of magnitude slower than on-chip bandwidth — minimizing host-device transfers is critical.
Within the GPU, the memory hierarchy goes: global VRAM → L2 cache → shared memory (per SM) → registers. High-performance CUDA kernels are written to maximize reuse in shared memory and avoid redundant global reads.
TPUs
Google's Tensor Processing Units (TPUs) are custom ASICs designed specifically for matrix multiplication workloads at bfloat16 precision. TPUs are organized into "pods" — up to thousands of chips connected by a high-bandwidth inter-chip interconnect (ICI). They excel at large batch training with regular tensor shapes and are tightly integrated with JAX and TensorFlow. For PyTorch users, torch_xla provides TPU support.
A kernel is compute-bound if arithmetic throughput is the bottleneck, and memory-bound if memory bandwidth is the bottleneck. Large matrix multiplications are typically compute-bound; elementwise operations (activation functions, layer norm) are memory-bound. Fused kernels combine multiple memory-bound ops into one kernel pass to cut memory traffic — this is the key idea behind FlashAttention.
Data Parallelism
The simplest and most common scaling strategy is data parallelism (DP): replicate the full model on each GPU, split the training batch across GPUs, compute gradients independently on each, then aggregate (average) the gradients before the parameter update.
After the backward pass, GPUs synchronize gradients via an All-Reduce collective operation — each GPU sends its gradients and receives the sum (or average) across all GPUs. Modern frameworks use the Ring-AllReduce algorithm, which distributes the communication load evenly and scales near-linearly in bandwidth efficiency.
PyTorch DDP
PyTorch's DistributedDataParallel (DDP) is the standard implementation. DDP overlaps gradient communication with the backward pass — as soon as a layer's gradients are computed, they begin synchronizing in the background while the backward pass continues through earlier layers. This hides most of the communication latency behind compute.
Linear Scaling Rule
When scaling from 1 GPU to N GPUs with data parallelism, the effective batch size increases N-fold. Empirically, the linear scaling rule holds: multiply the learning rate by N to maintain the same loss curve. However, this rule breaks down for very large batch sizes — the gradient noise becomes too small and the optimizer struggles to generalize. Warmup schedules help stabilize early training at large batch sizes.
Model Parallelism
When a model is too large to fit in a single GPU's VRAM, model parallelism is required: partition the model itself across multiple devices. The simplest form is pipeline parallelism, where different layers are assigned to different GPUs. The forward pass flows through GPU 1 (layers 1–N/K), then GPU 2 (layers N/K+1–2N/K), and so on.
Pipeline Bubbles
Naïve pipeline parallelism suffers from pipeline bubbles: GPUs downstream sit idle waiting for activations from upstream GPUs. The GPipe approach splits each batch into smaller micro-batches and pipelines them — GPU 2 processes micro-batch 1 while GPU 1 processes micro-batch 2, hiding much of the bubble overhead.
Tensor Parallelism
Tensor parallelism (also called intra-layer model parallelism) splits individual matrix multiplications across GPUs. For a weight matrix W and input X, the columns of W can be distributed across K GPUs; each GPU computes a partial result, then an All-Reduce collects the full output. This is the approach used by Megatron-LM for training large transformer models efficiently on NVLink-connected GPU clusters.
3D Parallelism
State-of-the-art large model training (GPT-3, LLaMA, Gemini) combines data parallelism + pipeline parallelism + tensor parallelism — so-called "3D parallelism." The optimal configuration depends on model size, GPU count, inter-GPU bandwidth (NVLink vs. InfiniBand vs. PCIe), and memory capacity.
ZeRO and FSDP
ZeRO (Zero Redundancy Optimizer, from DeepSpeed) identifies that standard data parallelism stores redundant copies of optimizer state, gradients, and parameters on every GPU. ZeRO eliminates this redundancy by sharding each across all GPUs:
- ZeRO-1 — shard optimizer state only (lowest communication overhead)
- ZeRO-2 — shard optimizer state + gradients
- ZeRO-3 — shard optimizer state + gradients + parameters (enables models far larger than any single GPU's VRAM)
PyTorch's native equivalent is FSDP (Fully Sharded Data Parallel), which implements a similar sharding strategy as ZeRO-3. FSDP is now the recommended approach for large model training in the PyTorch ecosystem.
Mixed Precision Training
Full-precision (FP32) training uses 4 bytes per number. Mixed precision training stores weights and gradients in half precision (FP16 or bfloat16 — 2 bytes each), reducing memory usage and doubling throughput on hardware with dedicated half-precision tensor cores (NVIDIA Ampere and Hopper have 2× more FP16/BF16 TFLOPS than FP32).
FP16 vs. BF16
| Format | Exponent bits | Mantissa bits | Dynamic range | Notes |
|---|---|---|---|---|
| FP32 | 8 | 23 | ~1.2×10⁻³⁸ to 3.4×10³⁸ | Standard; 4 bytes |
| FP16 | 5 | 10 | ~6×10⁻⁸ to 65504 | Loss scaling required to avoid underflow |
| BF16 | 8 | 7 | Same as FP32 | Same exponent range; preferred for stability |
BF16 has the same dynamic range as FP32 (8 exponent bits), making it far more numerically stable than FP16 for training. FP16 requires loss scaling — multiply the loss by a large constant before the backward pass, then divide gradients back — to prevent underflow in small gradient values. BF16 avoids this complication and is now the default choice on recent hardware.
Automatic Mixed Precision (AMP)
PyTorch's torch.cuda.amp.autocast automatically casts operations to the appropriate precision: matrix multiplications run in FP16/BF16, while reductions and numerically sensitive operations stay in FP32. The GradScaler handles loss scaling automatically for FP16 training.
Gradient Accumulation
Large batch sizes improve training stability and enable the linear scaling rule, but a large batch may not fit in GPU memory. Gradient accumulation simulates a large batch by splitting it into K smaller micro-batches processed sequentially. Gradients from each micro-batch are accumulated (summed) in-place. After K steps, the optimizer update runs once — exactly as if the full batch had been processed simultaneously.
Gradient accumulation is particularly useful when you want to match a published training configuration (which used many GPUs with large batches) on fewer GPUs. It comes at the cost of K times more forward+backward passes per optimizer step, but requires no extra memory beyond one micro-batch.
During the backward pass, activations from the forward pass must be held in memory to compute gradients. For deep networks this dominates VRAM usage. Gradient checkpointing (also called activation recomputation) discards activations during the forward pass and recomputes them on-the-fly during the backward pass — trading compute for memory. This typically doubles training time but can cut VRAM usage by 5–10× for deep transformers.
Distributed Training Frameworks
Several frameworks extend standard training loops to multi-GPU and multi-node settings:
Cloud ML Platforms
For most organizations, large-scale training happens in the cloud rather than on-premises. The major cloud providers offer managed ML training infrastructure:
- AWS SageMaker — managed training jobs with spot instance support; integrates with S3 for data; supports DDP, FSDP, and SageMaker-specific distributed libraries.
- Google Cloud Vertex AI — native TPU and GPU training; integration with Google Cloud Storage; supports JAX, TensorFlow, and PyTorch.
- Azure ML — managed compute clusters with RDMA networking for fast multi-node training; DeepSpeed integration.
- Lambda Labs / CoreWeave / Together AI — GPU cloud providers focused on ML workloads, often cheaper for burst training jobs.
Spot / Preemptible Instances
Cloud GPU instances are expensive. Spot instances (AWS) or preemptible VMs (GCP) offer the same hardware at 60–90% discount, but can be reclaimed by the provider at any time. Making training fault-tolerant — saving checkpoints frequently and resuming from the last checkpoint after preemption — is essential when using spot instances for long training runs.
GPUs accelerate neural network training through massive parallelism; TPUs offer specialized efficiency at bfloat16 precision. Data parallelism scales training to many GPUs by splitting batches and averaging gradients via All-Reduce. Model parallelism enables models too large for single-GPU memory using pipeline, tensor, or ZeRO/FSDP sharding. Mixed precision (BF16) halves memory and doubles throughput with minimal accuracy loss. Gradient accumulation simulates large batches on limited hardware. Frameworks like PyTorch DDP, FSDP, and DeepSpeed make distributed training accessible; cloud platforms provide on-demand GPU/TPU clusters.