Training at Scale
From a single GPU to thousands — the hardware, algorithms, and frameworks that make large-scale machine learning possible.
GPU & TPU
GPUs accelerate matrix multiplications 10–100× over CPUs. TPUs are Google's custom chips for BF16 tensor workloads, connected by high-bandwidth inter-chip interconnects.
Split the Batch
Replicate the full model on each GPU. Split the mini-batch. Each GPU computes gradients independently, then All-Reduce averages them before the optimizer step.
Overlapping Comms & Compute
DDP starts synchronizing gradients while the backward pass is still running — earlier-layer gradients travel over the network while later layers are still being differentiated. Communication latency is hidden behind compute.
When the Model Doesn't Fit
- Pipeline parallelism — assign layer groups to different GPUs; use micro-batches to hide bubbles
- Tensor parallelism — split weight matrices across GPUs; requires NVLink bandwidth
- 3D parallelism — combine all three strategies for LLM pre-training
Shard Everything Across GPUs
- ZeRO-1 — shard optimizer state only
- ZeRO-2 — shard optimizer state + gradients
- ZeRO-3 / FSDP — shard params + grads + optimizer state; enables trillion-parameter models
- ZeRO-Infinity — offload to CPU / NVMe when GPU VRAM is exhausted
BF16 vs. FP16
- FP16 — 2×memory savings; limited dynamic range; needs loss scaling
- BF16 — same dynamic range as FP32; no loss scaling needed; preferred on Ampere/Hopper
- Keep FP32 master weights for optimizer update step
- Use
torch.cuda.amp.autocastfor automatic mixed precision
Mixed Precision Breakdown
For a model with P parameters, the total memory per GPU in mixed precision training is:
Adam optimizer adds 2 FP32 tensors (1st + 2nd moment), accounting for 8P bytes.
Gradient Accumulation & Checkpointing
Accumulation: process K micro-batches, sum gradients, update once. Checkpointing: discard activations during forward, recompute on-the-fly during backward.
The Ecosystem
- PyTorch DDP — standard data parallelism; overlapped grad sync
- PyTorch FSDP — ZeRO-3-style sharding; native PyTorch 2.x
- DeepSpeed — ZeRO-1/2/3 + NVMe offload; widely used for LLM training
- Megatron-LM — 3D parallelism on NVLink clusters
- HF Accelerate — unified launcher: same script, 1 GPU or 1000
On-Demand Clusters
- AWS SageMaker — managed jobs, spot instances, S3 integration
- GCP Vertex AI — native TPU + GPU; JAX / TF / PyTorch
- Azure ML — RDMA networking, DeepSpeed integration
- Spot / preemptible GPUs: 60–90% cheaper; must checkpoint frequently and resume after preemption
Key Takeaways
- GPUs parallelize matrix ops; TPUs specialize at BF16 with pod-scale bandwidth
- Data parallelism: replicate model, split batch, All-Reduce gradients
- Pipeline + tensor parallelism + ZeRO/FSDP handle models beyond single-GPU VRAM
- BF16 mixed precision: 2× memory savings, same dynamic range as FP32
- Gradient accumulation simulates large batches; checkpointing saves VRAM at compute cost
- Cloud platforms offer on-demand GPU/TPU clusters with spot discounts