ML 101
M11 · L01
Deep Learning in Practice

Training at Scale

From a single GPU to thousands — the hardware, algorithms, and frameworks that make large-scale machine learning possible.

01 / 13
ML 101
M11 · L01
Hardware

GPU & TPU

GPU
Thousands of CUDA cores · VRAM
TPU
Custom ASIC · BF16 · Pod clusters

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.

02 / 13
ML 101
M11 · L01
Data Parallelism

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.

Averaged Gradient
\nabla_{\theta}\mathcal{L}=\frac{1}{N}\sum_{i=1}^{N}\nabla_{\theta}\mathcal{L}_i
03 / 13
ML 101
M11 · L01
PyTorch DDP

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.

Linear Scaling Rule
Scale learning rate by N when using N GPUs. Warmup schedules stabilize training at very large batch sizes.
04 / 13
ML 101
M11 · L01
Model Parallelism

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
05 / 13
ML 101
M11 · L01
ZeRO & FSDP

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
06 / 13
ML 101
M11 · L01
Mixed Precision

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.autocast for automatic mixed precision
07 / 13
ML 101
M11 · L01
Memory Budget

Mixed Precision Breakdown

For a model with P parameters, the total memory per GPU in mixed precision training is:

Memory (bytes)
16P\text{ bytes}\;=\;2P+2P+4P+8P

Adam optimizer adds 2 FP32 tensors (1st + 2nd moment), accounting for 8P bytes.

08 / 13
ML 101
M11 · L01
Memory Tricks

Gradient Accumulation & Checkpointing

Accum
Simulate large batch · no extra VRAM
Ckpt
Recompute activations · 5–10× less VRAM

Accumulation: process K micro-batches, sum gradients, update once. Checkpointing: discard activations during forward, recompute on-the-fly during backward.

09 / 13
ML 101
M11 · L01
Distributed Frameworks

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
10 / 13
ML 101
M11 · L01
Cloud ML

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
11 / 13
ML 101
Training at Scale

Check what scales

Four questions on parallelism, memory and mixed precision — a skimmer will miss them.

Question 1 of 0
Score 0/0

12 / 13
ML 101
Key Takeaways
Summary

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
13 / 13