Advanced Training
Distributed Training for LLMs â Scaling Beyond Single GPUs
Training modern LLMs requires distributing computation across hundreds or thousands of GPUs. This guide covers the fundamental parallelism strategies, memory optimization techniques, and frameworks that make large-scale training possible.
- Data Parallelism â Replicate model across GPUs, split data batches
- Tensor Parallelism â Split individual layers across multiple GPUs
- Pipeline Parallelism â Split model layers across GPU stages
No single GPU can hold a trillion-parameter model. Distribution is not optional â it is the foundation of modern AI.
Distributed Training for LLMs
Training modern LLMs requires distributing computation across hundreds or thousands of GPUs. A single NVIDIA A100 has 80GB of GPU memory â far insufficient for models with hundreds of billions of parameters. This tutorial covers the parallelism strategies, memory optimization techniques, and frameworks that enable large-scale LLM training.
Why Distributed Training is Necessary
Memory Requirements
The memory required to train a model depends on its parameter count, batch size, and optimizer state:
Scaling Laws and Training Time
Parallelism Strategies
Data Parallelism (DP)
The simplest form of distributed training â replicate the model on every GPU and split the data:
Limitations:
- Each GPU must hold the entire model (parameters + optimizer states + gradients)
- Communication cost scales with model size (all-reduce of gradients)
- Cannot train models larger than single GPU memory
ZeRO (Zero Redundancy Optimizer)
DeepSpeed ZeRO eliminates memory redundancy in data parallelism:
| ZeRO Stage | What is Partitioned | Memory Savings |
|---|---|---|
| Stage 1 | Optimizer states | 4x reduction |
| Stage 2 | Optimizer states + Gradients | 8x reduction |
| Stage 3 | Optimizer states + Gradients + Parameters | Nx reduction (N = GPUs) |
Fully Sharded Data Parallel (FSDP)
PyTorch's native implementation of ZeRO-3:
import torch
import torch.nn as nn
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import ShardingStrategy
model = nn.TransformerDecoder(
nn.TransformerDecoderLayer(d_model=4096, nhead=32),
num_layers=80
)
model = FSDP(
model,
sharding_strategy=ShardingStrategy.FULL_SHARD,
mixed_precision=MixedPrecision(
param_dtype=torch.bfloat16,
reduce_dtype=torch.bfloat16,
),
auto_wrap_policy=transformer_auto_wrap_policy,
device_id=torch.cuda.current_device(),
)
for batch in dataloader:
loss = model(batch)
loss.backward()
optimizer.step()
Tensor Parallelism (TP)
Split individual layers across multiple GPUs:
For a linear layer Y = XW where W is split column-wise:
Pipeline Parallelism (PP)
Split model layers across GPU stages:
3D Parallelism
Modern LLM training combines all three strategies:
Mixed-Precision Training
FP16 and BF16
| Format | Bits | Range | Precision | Use Case |
|---|---|---|---|---|
| FP32 | 32 | +/-3.4x10^38 | High | Master weights |
| FP16 | 16 | +/-65504 | Lower | Forward/backward pass |
| BF16 | 16 | +/-3.4x10^38 | Lower | Forward/backward pass |
| INT8 | 8 | +/-128 | Very low | Inference |
Loss Scaling
scaler = torch.cuda.amp.GradScaler()
for batch in dataloader:
optimizer.zero_grad()
with torch.cuda.amp.autocast(dtype=torch.bfloat16):
outputs = model(batch["input_ids"])
loss = criterion(outputs, batch["labels"])
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(optimizer)
scaler.update()
Communication Optimization
Gradient Accumulation
accumulation_steps = 8
optimizer.zero_grad()
for i, batch in enumerate(dataloader):
loss = model(batch) / accumulation_steps
loss.backward()
if (i + 1) % accumulation_steps == 0:
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
optimizer.zero_grad()
Communication Overhead
Training Frameworks
DeepSpeed
import deepspeed
ds_config = {
"train_batch_size": 1024,
"gradient_accumulation_steps": 8,
"fp16": {"enabled": True, "loss_scale": 0, "initial_scale_power": 16},
"zero_optimization": {
"stage": 3,
"overlap_comm": True,
"contiguous_gradients": True,
"reduce_bucket_size": 5e7,
"stage3_prefetch_bucket_size": 5e7,
"stage3_param_persistence_threshold": 1e6
},
"activation_checkpointing": {
"partition_activations": True,
"cpu_checkpointing": False,
"contiguous_memory_optimization": False
}
}
model_engine, optimizer, _, _ = deepspeed.initialize(
model=model,
config=ds_config,
model_parameters=model.parameters()
)
Activation Checkpointing
Memory Optimization Techniques
CPU Offloading
ds_config = {
"zero_optimization": {
"stage": 3,
"offload_optimizer": {"device": "cpu", "pin_memory": True},
"offload_param": {"device": "cpu", "pin_memory": True}
}
}
Training Stability
Learning Rate Scheduling
from transformers import get_cosine_schedule_with_warmup
scheduler = get_cosine_schedule_with_warmup(
optimizer,
num_warmup_steps=2000,
num_training_steps=100000,
min_lr_ratio=0.1
)
Gradient Clipping
max_grad_norm = 1.0
torch.nn.utils.clip_grad_norm_(model.parameters(), max_grad_norm)
Practice Exercises
-
ZeRO Stage Comparison: Calculate the memory savings for training a 7B model using ZeRO Stage 1, 2, and 3 on 4 GPUs. What is the memory per GPU for each stage?
-
Pipeline Bubble Analysis: Compute the pipeline efficiency for a 32-stage pipeline with 8, 16, 32, and 64 micro-batches. At what point does efficiency exceed 90%?
-
3D Parallelism Design: Design a 3D parallelism configuration for a 405B parameter model on 2048 GPUs. Justify your choice of TP, PP, and DP degrees.
-
Communication Analysis: For a 70B model with gradient all-reduce on 256 GPUs, calculate the communication volume and estimated latency assuming 100 Gbps interconnect.
Key Takeaways
What to Learn Next
-> Data Quality and Curation for LLMs How data quality, deduplication, and filtering impact model performance.
-> Knowledge Distillation for LLMs Compressing large models into smaller, faster ones without losing capability.
-> Flash Attention and Memory Efficiency IO-aware attention algorithms that reduce memory and increase speed.
-> Model Parallelism and Tensor Parallelism Deep dive into splitting models across GPUs for training and inference.
-> Training Deep Networks Foundational techniques for training neural networks effectively.
-> Scaling Laws and Chinchilla How model performance scales with compute, data, and parameters.