Home / ML 101 / Module 11 / Lesson 1

Training at Scale

GPU/TPU computing, data and model parallelism, mixed precision, and distributed training — the engineering that makes large-scale ML possible.

~16 min read M11 · L1 Advanced

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.

Compute vs. Memory Bound

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.

Data Parallel Gradient
\nabla_{\theta} \mathcal{L} = \frac{1}{N}\sum_{i=1}^{N} \nabla_{\theta} \mathcal{L}_i
Each GPU i computes gradients on its mini-batch shard. The averaged gradient is mathematically identical to computing on the full batch — provided learning rate is scaled appropriately.

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.

Data Parallelism
Replicate model, split data
Best for: models that fit in 1 GPU. Scales to hundreds of GPUs. Linear scaling rule for LR. All-Reduce after backward.
Pipeline Parallelism
Split layers across GPUs
Best for: models too large for 1 GPU. Use micro-batching to hide pipeline bubbles. Requires careful load balancing.
Tensor Parallelism
Split weight matrices
Best for: very large layers (attention, FFN). Requires high-bandwidth interconnect (NVLink). Used in Megatron-LM.
ZeRO / FSDP
Shard optimizer state
Shards optimizer state, gradients, and parameters across GPUs. PyTorch FSDP and DeepSpeed ZeRO-3 enable trillion-parameter training.

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:

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

Memory Savings
\text{Memory}_{\text{mixed}} = \underbrace{2P}_{\text{BF16 params}} + \underbrace{2P}_{\text{BF16 grads}} + \underbrace{4P}_{\text{FP32 master}} + \underbrace{8P}_{\text{Adam state}}
Parameters are stored in FP16/BF16 for forward/backward, but a master FP32 copy is maintained for the optimizer update to prevent precision loss in the weight update step.

FP16 vs. BF16

FormatExponent bitsMantissa bitsDynamic rangeNotes
FP32823~1.2×10⁻³⁸ to 3.4×10³⁸Standard; 4 bytes
FP16510~6×10⁻⁸ to 65504Loss scaling required to avoid underflow
BF1687Same as FP32Same 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.

Gradient Checkpointing

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:

PyTorch DDP
Data parallelism
Built-in, production-grade. Overlaps grad sync with backward pass. Launch with torchrun or SLURM srun.
PyTorch FSDP
Fully sharded DP
ZeRO-3-style sharding. Handles models larger than any single GPU. Native PyTorch 2.x support.
DeepSpeed
ZeRO optimizer
ZeRO-1/2/3 + ZeRO-Infinity (NVMe offload). Widely used for LLM pre-training. JSON config-driven.
Megatron-LM
3D parallelism
NVIDIA's framework combining data + pipeline + tensor parallelism. Used to train GPT-3-scale models.
Hugging Face Accelerate
Unified launcher
Abstracts DDP/FSDP/DeepSpeed behind a unified API. Launch the same training script on 1 GPU or 100.
JAX / XLA
Functional + TPU
jit + pmap/shard_map for data parallelism. XLA compiler fuses ops for TPUs and GPUs. Used by Google internally.

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:

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.


Key Takeaways

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.

Previous M10-L3: ML Pipeline Module Overview Next Lesson M11-L2: Transfer Learning